Finish 01_pytorch_workflow module.

This commit is contained in:
Emil
2026-07-05 04:36:10 +03:00
parent e86e768102
commit 53b0038a34
3 changed files with 354 additions and 104 deletions
@@ -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"
},
File diff suppressed because one or more lines are too long