Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions 01-tensor_tutorial.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -306,6 +306,14 @@
"v"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"> **Note**: `torch.Tensor(2, 3, 4)` creates an **uninitialized** tensor.\n",
"> For predictable values, prefer `torch.zeros`, `torch.ones`, or `torch.rand`.\n"
]
},
{
"cell_type": "code",
"execution_count": null,
Expand Down
3 changes: 2 additions & 1 deletion 02-space_stretching.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -113,7 +113,8 @@
" # transform points\n",
" Y = X @ W.t()\n",
" # compute singular values\n",
" U, S, V = torch.svd(W)\n",
" U, S, Vh = torch.linalg.svd(W)\n",
" V = Vh.T\n",
" # plot transformed points\n",
" show_scatterplot(Y, colors, title='y = Wx, singular values : [{:.3f}, {:.3f}]'.format(S[0], S[1]))\n",
" # transform the basis\n",
Expand Down
4 changes: 2 additions & 2 deletions 05-regression.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -173,8 +173,8 @@
"metadata": {},
"outputs": [],
"source": [
"plt.scatter(X.data.cpu().numpy(), y.data.cpu().numpy())\n",
"plt.plot(X.data.cpu().numpy(), y_pred.data.cpu().numpy(), 'r-', lw=5)\n",
"plt.scatter(X.detach().cpu().numpy(), y.detach().cpu().numpy())
"plt.plot(X.detach().cpu().numpy(), y_pred.detach().cpu().numpy(), 'r-', lw=5)
"plt.axis('equal');"
]
},
Expand Down
35 changes: 25 additions & 10 deletions 06-convnet.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,11 @@
"metadata": {},
"outputs": [],
"source": [
"from res.plot_lib import plot_data, plot_model, set_default"
"try:\n",
" from res.plot_lib import plot_data, plot_model, set_default\n",
" set_default()\n",
"except ImportError:\n",
" print(\"Warning: res.plot_lib not found, using default matplotlib settings.\")\n"
]
},
{
Expand All @@ -45,6 +49,11 @@
"set_default()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": []
},
{
"cell_type": "code",
"execution_count": null,
Expand Down Expand Up @@ -222,25 +231,31 @@
" model.eval()\n",
" test_loss = 0\n",
" correct = 0\n",
"\n",
" for data, target in test_loader:\n",
" # send to device\n",
" data, target = data.to(device), target.to(device)\n",
" \n",
"\n",
" # permute pixels\n",
" data = data.view(-1, 28*28)\n",
" data = data.view(-1, 28 * 28)\n",
" data = data[:, perm]\n",
" data = data.view(-1, 1, 28, 28)\n",
"\n",
" output = model(data)\n",
" test_loss += F.nll_loss(output, target, reduction='sum').item() # sum up batch loss \n",
" pred = output.data.max(1, keepdim=True)[1] # get the index of the max log-probability \n",
" test_loss += F.nll_loss(output, target, reduction='sum').item()\n",
" pred = output.data.max(1, keepdim=True)[1]\n",
" correct += pred.eq(target.data.view_as(pred)).cpu().sum().item()\n",
"\n",
" test_loss /= len(test_loader.dataset)\n",
" accuracy = 100. * correct / len(test_loader.dataset)\n",
" accuracy_list.append(accuracy)\n",
" print('\\nTest set: Average loss: {:.4f}, Accuracy: {}/{} ({:.0f}%)\\n'.format(\n",
" test_loss, correct, len(test_loader.dataset),\n",
" accuracy))"
" \n",
" print(\n",
" '\\nTest set: Average loss: {:.4f}, Accuracy: {}/{} ({:.0f}%)\\n'.format(\n",
" test_loss, correct, len(test_loader.dataset), accuracy\n",
" )\n",
" )\n",
"\n",
" return accuracy\n"
]
},
{
Expand All @@ -265,7 +280,7 @@
"\n",
"for epoch in range(0, 1):\n",
" train(epoch, model_fnn)\n",
" test(model_fnn)"
" acc = test(model_fnn)\n"
]
},
{
Expand Down