Finish 01_pytorch_workflow module.
This commit is contained in:
@@ -10,7 +10,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 266,
|
||||
"execution_count": 41,
|
||||
"id": "87dce210",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -25,7 +25,7 @@
|
||||
" 6: 'putting it all together'}"
|
||||
]
|
||||
},
|
||||
"execution_count": 266,
|
||||
"execution_count": 41,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -42,7 +42,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 267,
|
||||
"execution_count": 42,
|
||||
"id": "e90c8379",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -52,7 +52,7 @@
|
||||
"'2.12.1+cu130'"
|
||||
]
|
||||
},
|
||||
"execution_count": 267,
|
||||
"execution_count": 42,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -75,7 +75,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 268,
|
||||
"execution_count": 43,
|
||||
"id": "a5b9b36a",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -104,7 +104,7 @@
|
||||
" [0.4260]]))"
|
||||
]
|
||||
},
|
||||
"execution_count": 268,
|
||||
"execution_count": 43,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -126,7 +126,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 269,
|
||||
"execution_count": 44,
|
||||
"id": "41f90d81",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -136,7 +136,7 @@
|
||||
"(50, 50)"
|
||||
]
|
||||
},
|
||||
"execution_count": 269,
|
||||
"execution_count": 44,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -155,7 +155,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 270,
|
||||
"execution_count": 45,
|
||||
"id": "b2c8b4d9",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -165,7 +165,7 @@
|
||||
"(40, 40, 10, 10)"
|
||||
]
|
||||
},
|
||||
"execution_count": 270,
|
||||
"execution_count": 45,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -181,7 +181,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 271,
|
||||
"execution_count": 46,
|
||||
"id": "8cd57e5c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -227,7 +227,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 272,
|
||||
"execution_count": 47,
|
||||
"id": "aa9f92c3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
@@ -266,7 +266,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 273,
|
||||
"execution_count": 48,
|
||||
"id": "c97c753c",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -279,7 +279,7 @@
|
||||
" tensor([0.1288], requires_grad=True)]"
|
||||
]
|
||||
},
|
||||
"execution_count": 273,
|
||||
"execution_count": 48,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -297,7 +297,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 274,
|
||||
"execution_count": 49,
|
||||
"id": "f38d41c0",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -307,7 +307,7 @@
|
||||
"OrderedDict([('weights', tensor([0.3367])), ('bias', tensor([0.1288]))])"
|
||||
]
|
||||
},
|
||||
"execution_count": 274,
|
||||
"execution_count": 49,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -327,7 +327,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 275,
|
||||
"execution_count": 50,
|
||||
"id": "5263fef5",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -356,7 +356,7 @@
|
||||
" [0.9860]]))"
|
||||
]
|
||||
},
|
||||
"execution_count": 275,
|
||||
"execution_count": 50,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -367,7 +367,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 276,
|
||||
"execution_count": 51,
|
||||
"id": "01a46c08",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -386,7 +386,7 @@
|
||||
" [0.4588]])"
|
||||
]
|
||||
},
|
||||
"execution_count": 276,
|
||||
"execution_count": 51,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -400,7 +400,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 277,
|
||||
"execution_count": 52,
|
||||
"id": "e27f6f66",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -429,7 +429,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 278,
|
||||
"execution_count": 53,
|
||||
"id": "cfe9b6be",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -451,7 +451,7 @@
|
||||
" ))"
|
||||
]
|
||||
},
|
||||
"execution_count": 278,
|
||||
"execution_count": 53,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -486,7 +486,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 279,
|
||||
"execution_count": 54,
|
||||
"id": "c678f976",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -642,7 +642,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 280,
|
||||
"execution_count": 55,
|
||||
"id": "90cf11b7",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -801,7 +801,7 @@
|
||||
" 0.005023092031478882])"
|
||||
]
|
||||
},
|
||||
"execution_count": 280,
|
||||
"execution_count": 55,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -812,17 +812,17 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 281,
|
||||
"execution_count": 56,
|
||||
"id": "2d17ee65",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<matplotlib.legend.Legend at 0x7068be601e90>"
|
||||
"<matplotlib.legend.Legend at 0x730a131adc90>"
|
||||
]
|
||||
},
|
||||
"execution_count": 281,
|
||||
"execution_count": 56,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
},
|
||||
@@ -848,7 +848,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 282,
|
||||
"execution_count": 57,
|
||||
"id": "6aa42ea1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
@@ -859,7 +859,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 283,
|
||||
"execution_count": 58,
|
||||
"id": "04ce6868",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -869,7 +869,7 @@
|
||||
"OrderedDict([('weights', tensor([0.6990])), ('bias', tensor([0.3093]))])"
|
||||
]
|
||||
},
|
||||
"execution_count": 283,
|
||||
"execution_count": 58,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -881,7 +881,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 284,
|
||||
"execution_count": 59,
|
||||
"id": "f1a4db10",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -891,7 +891,7 @@
|
||||
"(0.7, 0.3)"
|
||||
]
|
||||
},
|
||||
"execution_count": 284,
|
||||
"execution_count": 59,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
@@ -902,7 +902,7 @@
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 285,
|
||||
"execution_count": 60,
|
||||
"id": "94cdacb6",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
@@ -921,10 +921,179 @@
|
||||
"plot_predictions(predictions=y_preds_new)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "70ce0795",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Saving a model in PyTorch"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 61,
|
||||
"id": "c4756ad1",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Saving model to: models/01_pytorch_workflow_model_0.pth\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"# Saving our PyTorch model\n",
|
||||
"from pathlib import Path\n",
|
||||
"\n",
|
||||
"# 1. Create models directory\n",
|
||||
"MODEL_PATH = Path(\"models\")\n",
|
||||
"MODEL_PATH.mkdir(parents=True, exist_ok=True)\n",
|
||||
"# 2. Create a model save path\n",
|
||||
"MODEL_NAME = \"01_pytorch_workflow_model_0.pth\"\n",
|
||||
"MODEL_SAVE_PATH = MODEL_PATH / MODEL_NAME\n",
|
||||
"MODEL_SAVE_PATH\n",
|
||||
"# 3. Save the model state_dict\n",
|
||||
"print(f\"Saving model to: {MODEL_SAVE_PATH}\")\n",
|
||||
"torch.save(obj=model_0.state_dict(), f=MODEL_SAVE_PATH)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 62,
|
||||
"id": "31c5e50d",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"total 4\n",
|
||||
"-rw-rw-r-- 1 emil emil 2117 Jul 5 03:40 01_pytorch_workflow_model_0.pth\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"!ls -l models"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "18a405a2",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Loading a PyTorch model"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 63,
|
||||
"id": "ef3cf06a",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<All keys matched successfully>"
|
||||
]
|
||||
},
|
||||
"execution_count": 63,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"# To load in a saved state_dict we have to instantiane a new instance of our model class\n",
|
||||
"\n",
|
||||
"loaded_model_0 = LinearRegressionModel()\n",
|
||||
"\n",
|
||||
"# Load the saved state_dict of model_0\n",
|
||||
"loaded_model_0.load_state_dict(torch.load(f=MODEL_SAVE_PATH))\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 64,
|
||||
"id": "82b3e972",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([[0.8685],\n",
|
||||
" [0.8825],\n",
|
||||
" [0.8965],\n",
|
||||
" [0.9105],\n",
|
||||
" [0.9245],\n",
|
||||
" [0.9384],\n",
|
||||
" [0.9524],\n",
|
||||
" [0.9664],\n",
|
||||
" [0.9804],\n",
|
||||
" [0.9944]])"
|
||||
]
|
||||
},
|
||||
"execution_count": 64,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"# Make some preditions with our loaded model\n",
|
||||
"loaded_model_0.eval()\n",
|
||||
"with torch.inference_mode():\n",
|
||||
" loaded_model_preds = loaded_model_0(X_test)\n",
|
||||
"\n",
|
||||
"loaded_model_preds"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 66,
|
||||
"id": "fb05340b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"model_0.eval()\n",
|
||||
"with torch.inference_mode():\n",
|
||||
" y_preds = model_0(X_test)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 67,
|
||||
"id": "41a4ed50",
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"tensor([[True],\n",
|
||||
" [True],\n",
|
||||
" [True],\n",
|
||||
" [True],\n",
|
||||
" [True],\n",
|
||||
" [True],\n",
|
||||
" [True],\n",
|
||||
" [True],\n",
|
||||
" [True],\n",
|
||||
" [True]])"
|
||||
]
|
||||
},
|
||||
"execution_count": 67,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"# Compare loaded model preds with the original model preds\n",
|
||||
"y_preds == loaded_model_preds"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "70ce0795",
|
||||
"id": "da9c5c67",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
@@ -932,7 +1101,7 @@
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": ".venv (3.11.15.final.0)",
|
||||
"display_name": ".venv (3.11.15)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
|
||||
Binary file not shown.
+148
-67
File diff suppressed because one or more lines are too long
Reference in New Issue
Block a user