tahamajs's picture
download
raw
24 kB
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [],
"source": [
"import torch\n",
"import torch.nn as nn\n",
"import torch.optim as optim\n",
"import torch.nn.functional as F\n",
"from torchvision.datasets import MNIST\n",
"from torchvision.transforms import transforms\n",
"from torch.utils.data import DataLoader\n",
"import matplotlib.pyplot as plt\n",
"from torchvision.transforms import ToTensor"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Set Device"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [],
"source": [
"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Load MNIST"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [],
"source": [
"MNIST_MEAN = 0.1307\n",
"MNIST_STD = 0.3081\n",
"\n",
"transform = transforms.Compose([\n",
" transforms.ToTensor(), \n",
" transforms.Normalize((MNIST_MEAN,), (MNIST_STD,))\n",
" ])\n",
"\n",
"train_dataset = MNIST(root='../data', train=True, download=True, transform=transform)\n",
"test_dataset = MNIST(root='../data', train=False, download=True, transform=transform)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# MLP"
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [],
"source": [
"class MLP(nn.Module):\n",
" def __init__(self, input_shape, hidden_size, num_classes):\n",
" super(MLP, self).__init__()\n",
" self.input_shape = input_shape\n",
" self.fc1 = nn.Linear(input_shape[1], hidden_size)\n",
" self.relu1 = nn.ReLU()\n",
" self.fc2 = nn.Linear(hidden_size, num_classes)\n",
" self.softmax = nn.Softmax(dim=1)\n",
"\n",
" def forward(self, x):\n",
" out = self.fc1(x)\n",
" out = self.relu1(out)\n",
" out = self.fc2(out)\n",
" out = self.softmax(out)\n",
" return out"
]
},
{
"cell_type": "code",
"execution_count": 7,
"metadata": {},
"outputs": [],
"source": [
"class TrainTest:\n",
" def __init__(self, model, device) -> None:\n",
" self.model = model\n",
" self.device = device\n",
" \n",
" def train(self, num_epochs, train_dataset, batch_size, loss_function, optimizer):\n",
" train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\n",
" \n",
" for epoch in range(num_epochs):\n",
" for i, (images, labels) in enumerate(train_loader):\n",
" images = images.reshape(*self.model.input_shape).to(self.device)\n",
" labels = labels.to(self.device)\n",
"\n",
" outputs = self.model(images)\n",
" loss = loss_function(outputs, labels)\n",
"\n",
" optimizer.zero_grad()\n",
" loss.backward()\n",
" optimizer.step()\n",
"\n",
" if (i + 1) % 100 == 0:\n",
" self._print_log(epoch, num_epochs, i + 1, len(train_loader), loss.item())\n",
" \n",
" def test(self, test_dataset, batch_size):\n",
" test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)\n",
" \n",
" with torch.no_grad():\n",
" correct = 0\n",
" total = 0\n",
" for images, labels in test_loader:\n",
" images = images.reshape(*self.model.input_shape).to(self.device)\n",
" labels = labels.to(self.device)\n",
" outputs = self.model(images)\n",
" _, predicted = torch.max(outputs.data, 1)\n",
" total += labels.size(0)\n",
" correct += (predicted == labels).sum().item()\n",
"\n",
" print(f'Accuracy on test dataset: {100 * correct / total}%')\n",
" \n",
" def _print_log(self, epoch, num_epochs, step, total_steps, loss):\n",
" print(f'Epoch [{epoch + 1}/{num_epochs}], Step [{step}/{total_steps}], Loss: {loss:.4f}')\n",
" \n",
" def predict(self, x):\n",
" x = x.reshape(*self.model.input_shape).to(self.device)\n",
" outputs = self.model(x)\n",
" _, predicted = torch.max(outputs.data, 1)\n",
" return predicted\n",
" "
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch [1/2], Step [100/938], Loss: 1.7355\n",
"Epoch [1/2], Step [200/938], Loss: 1.6238\n",
"Epoch [1/2], Step [300/938], Loss: 1.5739\n",
"Epoch [1/2], Step [400/938], Loss: 1.5461\n",
"Epoch [1/2], Step [500/938], Loss: 1.5804\n",
"Epoch [1/2], Step [600/938], Loss: 1.5567\n",
"Epoch [1/2], Step [700/938], Loss: 1.5231\n",
"Epoch [1/2], Step [800/938], Loss: 1.5747\n",
"Epoch [1/2], Step [900/938], Loss: 1.5845\n",
"Epoch [2/2], Step [100/938], Loss: 1.5401\n",
"Epoch [2/2], Step [200/938], Loss: 1.5773\n",
"Epoch [2/2], Step [300/938], Loss: 1.5643\n",
"Epoch [2/2], Step [400/938], Loss: 1.5667\n",
"Epoch [2/2], Step [500/938], Loss: 1.5499\n",
"Epoch [2/2], Step [600/938], Loss: 1.4885\n",
"Epoch [2/2], Step [700/938], Loss: 1.5458\n",
"Epoch [2/2], Step [800/938], Loss: 1.5290\n",
"Epoch [2/2], Step [900/938], Loss: 1.5581\n",
"Accuracy on test dataset: 93.97%\n"
]
}
],
"source": [
"hidden_size = 50\n",
"num_classes = 10\n",
"num_epochs = 2\n",
"batch_size = 64\n",
"learning_rate = 0.001\n",
"input_shape = (-1, 784)\n",
"\n",
"mlp_clf = MLP(input_shape, hidden_size, num_classes).to(device)\n",
"\n",
"loss_function = nn.CrossEntropyLoss()\n",
"optimizer = optim.Adam(mlp_clf.parameters(), lr=learning_rate)\n",
"\n",
"TrainTest(mlp_clf, device).train(num_epochs, train_dataset, batch_size, loss_function, optimizer)\n",
"\n",
"TrainTest(mlp_clf, device).test(test_dataset, batch_size)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# LeNet"
]
},
{
"cell_type": "code",
"execution_count": 13,
"metadata": {},
"outputs": [],
"source": [
"class LeNet(nn.Module):\n",
" def __init__(self, input_shape):\n",
" super(LeNet, self).__init__()\n",
" self.input_shape = input_shape\n",
" self.conv1 = nn.Conv2d(1, 6, 5, padding=2)\n",
" self.pool1 = nn.AvgPool2d(2, stride=2)\n",
" self.conv2 = nn.Conv2d(6, 16, 5)\n",
" self.pool2 = nn.AvgPool2d(2, 2)\n",
" self.fc1 = nn.Linear(16 * 5 * 5, 120)\n",
" self.fc2 = nn.Linear(120, 84)\n",
" self.fc3 = nn.Linear(84, 10)\n",
"\n",
" def forward(self, x):\n",
" x = F.sigmoid(self.conv1(x))\n",
" x = self.pool1(x)\n",
" x = F.sigmoid(self.conv2(x))\n",
" x = self.pool2(x)\n",
" x = x.view(-1, 16 * 5 * 5)\n",
" x = F.sigmoid(self.fc1(x))\n",
" x = F.sigmoid(self.fc2(x))\n",
" x = self.fc3(x)\n",
" return x"
]
},
{
"cell_type": "code",
"execution_count": 15,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch [1/4], Step [100/938], Loss: 2.3060\n",
"Epoch [1/4], Step [200/938], Loss: 1.5703\n",
"Epoch [1/4], Step [300/938], Loss: 0.9419\n",
"Epoch [1/4], Step [400/938], Loss: 0.5567\n",
"Epoch [1/4], Step [500/938], Loss: 0.4218\n",
"Epoch [1/4], Step [600/938], Loss: 0.4604\n",
"Epoch [1/4], Step [700/938], Loss: 0.2144\n",
"Epoch [1/4], Step [800/938], Loss: 0.3131\n",
"Epoch [1/4], Step [900/938], Loss: 0.3399\n",
"Epoch [2/4], Step [100/938], Loss: 0.2391\n",
"Epoch [2/4], Step [200/938], Loss: 0.2738\n",
"Epoch [2/4], Step [300/938], Loss: 0.1024\n",
"Epoch [2/4], Step [400/938], Loss: 0.1431\n",
"Epoch [2/4], Step [500/938], Loss: 0.2770\n",
"Epoch [2/4], Step [600/938], Loss: 0.1894\n",
"Epoch [2/4], Step [700/938], Loss: 0.2015\n",
"Epoch [2/4], Step [800/938], Loss: 0.0950\n",
"Epoch [2/4], Step [900/938], Loss: 0.1920\n",
"Epoch [3/4], Step [100/938], Loss: 0.1218\n",
"Epoch [3/4], Step [200/938], Loss: 0.1982\n",
"Epoch [3/4], Step [300/938], Loss: 0.2418\n",
"Epoch [3/4], Step [400/938], Loss: 0.1151\n",
"Epoch [3/4], Step [500/938], Loss: 0.0518\n",
"Epoch [3/4], Step [600/938], Loss: 0.1406\n",
"Epoch [3/4], Step [700/938], Loss: 0.1496\n",
"Epoch [3/4], Step [800/938], Loss: 0.3267\n",
"Epoch [3/4], Step [900/938], Loss: 0.1832\n",
"Epoch [4/4], Step [100/938], Loss: 0.0407\n",
"Epoch [4/4], Step [200/938], Loss: 0.0760\n",
"Epoch [4/4], Step [300/938], Loss: 0.2648\n",
"Epoch [4/4], Step [400/938], Loss: 0.1856\n",
"Epoch [4/4], Step [500/938], Loss: 0.1392\n",
"Epoch [4/4], Step [600/938], Loss: 0.0647\n",
"Epoch [4/4], Step [700/938], Loss: 0.0561\n",
"Epoch [4/4], Step [800/938], Loss: 0.1096\n",
"Epoch [4/4], Step [900/938], Loss: 0.1033\n",
"Accuracy on test dataset: 97.64%\n"
]
}
],
"source": [
"num_epochs = 4\n",
"input_shape = (-1, 1, 28, 28)\n",
"\n",
"leNet_clf = LeNet(input_shape).to(device)\n",
"\n",
"loss_function = nn.CrossEntropyLoss()\n",
"optimizer = optim.Adam(leNet_clf.parameters(), lr=learning_rate)\n",
"\n",
"TrainTest(leNet_clf, device).train(num_epochs, train_dataset, batch_size, loss_function, optimizer)\n",
"\n",
"TrainTest(leNet_clf, device).test(test_dataset, batch_size)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Test Shifted Image"
]
},
{
"cell_type": "code",
"execution_count": 37,
"metadata": {},
"outputs": [],
"source": [
"def show_image(image, ax, title):\n",
" ax.imshow(image.squeeze(), cmap='gray')\n",
" ax.set_title(title)\n",
" ax.axis('off')"
]
},
{
"cell_type": "code",
"execution_count": 40,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
" original shifted\n",
"LeNet 7 7\n",
"MLP 7 2\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAxoAAAGICAYAAADGcZYzAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjcuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8pXeV/AAAACXBIWXMAAA9hAAAPYQGoP6dpAAAg2ElEQVR4nO3da5BV1Z3w4f+Bbq6NhItKKyAIUZkMDCAaoFGBMIqgMSJeYhwvRRAVhZRIxKoAijoQlKAxAa2J4szEazIODkghKGouQhJMxIEwGmBQEfFCm4hE5LbfD7z0pAWEjattoJ+nyg/sXuustU9ZZ/Prfc6hkGVZFgAAAAnVqu4NAAAAhx6hAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhUcMtWrQozj///CgtLY06depEixYtYvDgwbFw4cJcj3PzzTdHoVDYrz08//zzUSgU4vnnn9+v+fuqd+/e0bt3730a9/d///dVuhcA9u43v/lNnHvuudG6deuoW7duHHnkkdGjR48YNWpUpXFt2rSJs846a6+Pt6frzT333BPt27ePOnXqRKFQiD//+c/xz//8zzFz5syEZ7PD6tWro1AoxIMPPrhPe/35z3+efA/wRREaNdg999wTZWVlsWbNmpg8eXI888wzceedd8Zbb70VvXr1ih/96Ef7/Fjf/va3c8fJTl27do2FCxdG165d92s+AIeep556Knr27BkffvhhTJ48OebNmxd33313lJWVxWOPPbZfj7m7683LL78cI0aMiD59+sSCBQti4cKF0ahRoyoLDahJiqp7A1SPX//61/Gd73wnBgwYEP/5n/8ZRUX/97/CRRddFOeee26MHDkyunTpEmVlZXt8nL/+9a/RoEGDaNmyZbRs2XK/9nLYYYdF9+7d92suAIemyZMnR9u2bePpp5/e5Ro1efLk/XrM3V1vli1bFhERQ4cOjZNPPnn/Nwzswh2NGmrixIlRKBRi+vTplV7AIyKKiopi2rRpUSgUYtKkSRXHd7496ve//30MHjw4mjRpEu3atav0s7/1ySefxKhRo6JFixbRoEGDOPXUU+Oll16KNm3axOWXX14xbne3si+//PIoKSmJFStWxIABA6KkpCRatWoVo0aNik8++aTSOrfcckt89atfjaZNm8Zhhx0WXbt2jfvvvz+yLEv0bEUUCoW49tprY8aMGXH88cdH/fr1o1u3brFo0aLIsizuuOOOaNu2bZSUlETfvn1jxYoVlebPnz8/zjnnnGjZsmXUq1cv2rdvH8OGDYv3339/l7WefPLJ6NSpU9StWzeOPfbYuPvuu3f7/GZZFtOmTYvOnTtH/fr1o0mTJjF48OBYtWpVsvMGqC7r16+P5s2b73KNioioVWv3f32ZO3dudO3aNerXrx8nnHBCPPDAA5V+/unrTe/eveOSSy6JiIivfvWrUSgU4vLLL49CoRAbN26Mf/3Xf41CoRCFQqHSW2/XrVsXw4YNi5YtW0adOnWibdu2ccstt8TWrVsrrbd27dq44IILolGjRtG4ceO48MILY926dfv9nOy8Frzyyitx/vnnR+PGjaNp06Zx/fXXx9atW+PVV1+N/v37R6NGjaJNmza7BNmmTZti1KhR0blz54q5PXr0iCeffHKXtf785z/HkCFDomnTplFSUhIDBw6MVatWRaFQiJtvvrnS2D/96U9x8cUXxxFHHBF169aNDh06xI9//OP9Pk8OHe5o1EDbtm2L5557Lrp167bHuxCtWrWKE088MRYsWBDbtm2L2rVrV/xs0KBBcdFFF8VVV10VGzdu3OM6V1xxRTz22GPx3e9+N/r27Rt//OMf49xzz40PP/xwn/a5ZcuW+PrXvx5DhgyJUaNGxS9+8Yu49dZbo3HjxjFu3LiKcatXr45hw4ZF69atI2LH506uu+66eOuttyqN+7xmz54df/jDH2LSpElRKBTixhtvjIEDB8Zll10Wq1atih/96Efxl7/8Ja6//vo477zz4uWXX66Ig5UrV0aPHj3i29/+djRu3DhWr14dP/jBD6JXr17x3//931FcXBwROy6SgwYNilNPPTUee+yx2Lp1a9x5553xzjvv7LKfYcOGxYMPPhgjRoyI73//+1FeXh4TJkyInj17xpIlS+LII49Mdu4AX7QePXrET37ykxgxYkR861vfiq5du1a8Vu7OkiVLYtSoUTFmzJg48sgj4yc/+UkMGTIk2rdvH6eeeupu50ybNi0eeeSRuO2222LGjBlxwgknxOGHHx5XXXVV9O3bN/r06RNjx46NiB13QyJ2RMbJJ58ctWrVinHjxkW7du1i4cKFcdttt8Xq1atjxowZERHx8ccfR79+/WLt2rUxceLEOO644+Kpp56KCy+88HM/NxdccEFccsklMWzYsJg/f35Mnjw5tmzZEs8880xcc801ccMNN8TDDz8cN954Y7Rv3z4GDRoUETt+AVheXh433HBDHH300bF58+Z45plnYtCgQTFjxoy49NJLIyJi+/btcfbZZ8fixYvj5ptvrnjLWf/+/XfZyx//+Mfo2bNntG7dOqZMmRItWrSIp59+OkaMGBHvv/9+jB8//nOfLwexjBpn3bp1WURkF1100WeOu/DCC7OIyN55550sy7Js/PjxWURk48aN22Xszp/ttGzZsiwishtvvLHSuEceeSSLiOyyyy6rOPbcc89lEZE999xzFccuu+yyLCKyxx9/vNL8AQMGZMcff/we97xt27Zsy5Yt2YQJE7JmzZpl27dvr/jZaaedlp122mmfec47x33lK1+pdCwishYtWmQfffRRxbGZM2dmEZF17ty50jp33XVXFhHZK6+8stvH3759e7Zly5bs9ddfzyIie/LJJyt+dtJJJ2WtWrXKPvnkk4pjGzZsyJo1a1bp+V24cGEWEdmUKVMqPfabb76Z1a9fP/vud7+71/MEOJC9//77Wa9evbKIyCIiKy4uznr27JlNnDgx27BhQ6WxxxxzTFavXr3s9ddfrzj28ccfZ02bNs2GDRtWcWx315sZM2ZkEZH97ne/q/SYDRs2rHSt2mnYsGFZSUlJpbWyLMvuvPPOLCKyZcuWZVmWZdOnT9/lNT7Lsmzo0KFZRGQzZsz4zPPfudef/exnFcd2Xms//drfuXPnLCKyJ554ouLYli1bssMPPzwbNGjQHtfYunVrtmXLlmzIkCFZly5dKo4/9dRTWURk06dPrzR+4sSJWURk48ePrzh2xhlnZC1btsz+8pe/VBp77bXXZvXq1cvKy8s/8zw5tHnrFHuU/f+3Hn36LTvnnXfeXue+8MILEbHjty5/a/Dgwbu9Db47hUIhzj777ErHOnXqFK+//nqlYwsWLIh+/fpF48aNo3bt2lFcXBzjxo2L9evXx7vvvrtPa+2LPn36RMOGDSv+3KFDh4iIOPPMMys9RzuP/+0+33333bjqqquiVatWUVRUFMXFxXHMMcdERMTy5csjImLjxo2xePHi+MY3vhF16tSpmFtSUrLL8zB79uwoFApxySWXxNatWyv+a9GiRfzDP/xDlX+DF0BVa9asWfzyl7+M3/3udzFp0qQ455xz4rXXXoubbropOnbsuMtbTzt37lxxZzsiol69enHcccftcs34vGbPnh19+vSJo446qtLr75lnnhkR/3f9e+6556JRo0bx9a9/vdL8iy+++HPv4dPfsNWhQ4coFAoVe4jY8Tbo9u3b73L+P/vZz6KsrCxKSkoqrkf3339/xbXob8/h09fwb37zm5X+vGnTpnj22Wfj3HPPjQYNGlR6PgYMGBCbNm2KRYsWfe7z5eDlrVM1UPPmzaNBgwbxv//7v585bvXq1dGgQYNo2rRppeOlpaV7XWP9+vUREbu8faeoqCiaNWu2T/ts0KBB1KtXr9KxunXrxqZNmyr+/Nvf/jZOP/306N27d/zLv/xLxftlZ86cGbfffnt8/PHH+7TWvvj087AzBvZ0fOc+t2/fHqeffnqsXbs2xo4dGx07doyGDRvG9u3bo3v37hV7/OCDDyLLst2+5enTx9555509jo2IOPbYY/fjDAEOPN26dYtu3bpFxI631N54440xderUmDx5cqXPIOzu2lK3bt2k14GIHa+/s2bN2uPbuHYG0Pr163f7Gt2iRYvPvYfdXXd2d82sU6dOpbcrP/HEE3HBBRfE+eefH6NHj44WLVpEUVFRTJ8+vdLnWdavXx9FRUW7rPPp81m/fn1s3bo17rnnnrjnnnt2u9fdfRaRmkNo1EC1a9eOPn36xNy5c2PNmjW7/ZzGmjVr4qWXXoozzzyz0uczIna9w7E7O1/w33nnnTj66KMrjm/durUiQlJ49NFHo7i4OGbPnl3pBfZA+krCpUuXxpIlS+LBBx+Myy67rOL4pz8w3qRJkygUCrv9PManPzzYvHnzKBQK8ctf/jLq1q27y/jdHQM42BUXF8f48eNj6tSpsXTp0mrZQ/PmzaNTp05x++237/bnRx11VETsuA7+9re/3eXnn+fD4J/XT3/602jbtm089thjla7ln/6SlWbNmsXWrVujvLy8Umx8eu9NmjSJ2rVrxz/90z/F8OHDd7tm27ZtE54BBxtvnaqhbrrppsiyLK655prYtm1bpZ9t27Ytrr766siyLG666ab9evydH7z79Hed//znP9/lWzk+j0KhEEVFRZVi6OOPP45///d/T7bG57XzxfzTf/m/7777Kv25YcOG0a1bt5g5c2Zs3ry54vhHH30Us2fPrjT2rLPOiizL4q233qr4bd/f/texY8cqOhuAL8bbb7+92+M73+Kz8y/0VWVPd0POOuusWLp0abRr1263r78799WnT5/YsGFD/Nd//Vel+Q8//HCV7vuzFAqFin+UcKd169bt8q1Tp512WkTseg1/9NFHK/25QYMG0adPn/jDH/4QnTp12u3zsa/vYuDQ5I5GDVVWVhZ33XVXfOc734levXrFtddeG61bt4433ngjfvzjH8dvfvObuOuuu6Jnz5779fhf+cpX4pvf/GZMmTIlateuHX379o1ly5bFlClTonHjxnv8asK8Bg4cGD/4wQ/i4osvjiuvvDLWr18fd9555wH1G/0TTjgh2rVrF2PGjIksy6Jp06Yxa9asmD9//i5jJ0yYEAMHDowzzjgjRo4cGdu2bYs77rgjSkpKory8vGJcWVlZXHnllXHFFVfE4sWL49RTT42GDRvG22+/Hb/61a+iY8eOcfXVV3+RpwmQ1BlnnBEtW7aMs88+O0444YTYvn17vPzyyzFlypQoKSmJkSNHVun6HTt2jOeffz5mzZoVpaWl0ahRozj++ONjwoQJMX/+/OjZs2eMGDEijj/++Ni0aVOsXr065syZE/fee2+0bNkyLr300pg6dWpceumlcfvtt8eXv/zlmDNnTjz99NNVuu/PctZZZ8UTTzwR11xzTQwePDjefPPNuPXWW6O0tDT+9Kc/VYzr379/lJWVxahRo+LDDz+ME088MRYuXBj/9m//FhGVv1747rvvjl69esUpp5wSV199dbRp0yY2bNgQK1asiFmzZsWCBQu+8PPkwCE0arDrrrsuTjrppJgyZUqMGjUq1q9fH02bNo1evXrFr371q+jRo8fnevwZM2ZEaWlp3H///TF16tTo3LlzPP7449G/f//40pe+lOQc+vbtGw888EB8//vfj7PPPjuOPvroGDp0aBxxxBExZMiQJGt8XsXFxTFr1qwYOXJkDBs2LIqKiqJfv37xzDPPVPrgYsSOF/f/+I//iHHjxsWFF14YLVq0iGuuuSbWrl27y12a++67L7p37x733XdfTJs2LbZv3x5HHXVUlJWV+UengIPe9773vXjyySdj6tSp8fbbb8cnn3wSpaWl0a9fv7jpppsqvnijqtx9990xfPjwuOiii+Kvf/1rnHbaafH8889HaWlpLF68OG699da44447Ys2aNdGoUaNo27Zt9O/fP5o0aRIRO37bv2DBghg5cmSMGTMmCoVCnH766fHoo4/u9y/xPq8rrrgi3n333bj33nvjgQceiGOPPTbGjBkTa9asiVtuuaViXK1atWLWrFkxatSomDRpUmzevDnKysripz/9aXTv3r3SNfzv/u7v4ve//33ceuut8b3vfS/efffd+NKXvhRf/vKXY8CAAdVwlhxIClmW8F81g7148cUXo6ysLB566KEk37xRE2zZsiU6d+4cRx99dMybN6+6twNADfXwww/Ht771rfj1r39dbbHEwUVoUGXmz58fCxcujBNPPDHq168fS5YsiUmTJkXjxo3jlVde2eXbMdhhyJAh8Y//+I9RWloa69ati3vvvTdeeOGFmDdvXvTr16+6twdADfDII4/EW2+9FR07doxatWrFokWL4o477oguXbpUfP0t7I23TlFlDjvssJg3b17cddddsWHDhmjevHmceeaZMXHiRJHxGTZs2BA33HBDvPfee1FcXBxdu3aNOXPmiAwAvjCNGjWKRx99NG677bbYuHFjlJaWxuWXXx633XZbdW+Ng4g7GgAAQHK+3hYAAEhOaAAAAMkJDQAAIDmhAQAAJLfP3zr1t/9cPQBfLN/bsXuuTQDVZ2/XJnc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJFdU3RsAAA4dgwcPzjV+6NChuddYu3Zt7jmbNm3KNf6hhx7Kvca6detyjV+xYkXuNeBg4o4GAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEiukGVZtk8DC4Wq3gsAe7CPL9U1jmvTgWfVqlW5xrdp06ZqNlINNmzYkGv8smXLqmgn7K81a9bknjN58uTccxYvXpx7zoFob9cmdzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJF1b0BAODQMXTo0FzjO3XqlHuN5cuX557ToUOHXOO7du2ae43evXvnGt+9e/fca7z55pu557Rq1Sr3nC/C1q1bc41/7733cq9RWlqae05eb7zxRu45ixcvroKdHHjc0QAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASK6QZVm2TwMLhareyyFl8ODBucYPHTo09xpr167NNX7Tpk2513jooYdyz1m3bl2u8StWrMi9BtQ0+/hSXeO4NnEgadKkSa7xnTt3zr3GSy+9lHvOSSedlHvOFyHv30tee+213GssX7481/imTZvmXmP48OG550yfPj33nAPR3q5N7mgAAADJCQ0AACA5oQEAACQnNAAAgOSEBgAAkJzQAAAAkhMaAABAckIDAABITmgAAADJCQ0AACA5oQEAACQnNAAAgOQKWZZl+zSwUKjqvRxSVq1alWt8mzZtqmYj1WDDhg25xi9btqyKdsL+WrNmTe45kydPzjV+8eLFudeoyfbxpbrGcW2CmuO8887LPefxxx/PNX7p0qW51+jTp0/uOeXl5bnnHIj2dm1yRwMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAILmi6t7AoWro0KG5xnfq1Cn3GsuXL881vkOHDrnX6Nq1a+45vXv3zjW+e/fuudd48803c41v1apV7jW+CFu3bs0957333ss9p7S0NPecvN54441c4xcvXlxFOwHgYHDEEUfkGj9t2rTca9Sqle936hMmTMi9Rnl5ee45NYU7GgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAguaLq3sCh6tlnn63S8ftj7ty5Vb5GRESTJk1yje/cuXPuNV566aVc40866aTca3wRNm3alHvOa6+9lnvO8uXLc41v2rRp7jVWrlyZew4ANdfw4cNzjT/88MNzr/HBBx/kGv/qq6/mXoM9c0cDAABITmgAAADJCQ0AACA5oQEAACQnNAAAgOSEBgAAkJzQAAAAkhMaAABAckIDAABITmgAAADJCQ0AACA5oQEAACRXyLIs26eBhUJV7wXYB+edd17uOY8//niu8UuXLs29Rp8+fXKNLy8vz71GTbaPL9U1jmsTHBjKyspyz1mwYEGu8cXFxbnX6N27d67xv/jFL3KvUZPt7drkjgYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJF1b0BqMmOOOKI3HOmTZuWe06tWvl+pzBhwoTca5SXl+eeA8ChYcCAAbnnFBcX5xr/7LPP5l5j4cKFueeQjjsaAABAckIDAABITmgAAADJCQ0AACA5oQEAACQnNAAAgOSEBgAAkJzQAAAAkhMaAABAckIDAABITmgAAADJCQ0AACC5oureANRkw4cPzz3n8MMPzz3ngw8+yDX+1Vdfzb0GAIeG+vXr557Tv3//3HM2b96ca/z48eNzr7Fly5bcc0jHHQ0AACA5oQEAACQnNAAAgOSEBgAAkJzQAAAAkhMaAABAckIDAABITmgAAADJCQ0AACA5oQEAACQnNAAAgOSKqnsDcCgpKyvLNX7MmDFVtJPKvvGNb+Qav3Tp0qrZCAAHvNGjR+ee06VLl9xz5s6dm2v8iy++mHsNqpc7GgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAguaLq3gAcSgYMGJBrfHFxce41nn322dxzFi5cmHsOAIeGgQMH5ho/duzY3Gt8+OGHuedMmDAh9xwOLu5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkiqp7A3Cgql+/fu45/fv3zzV+8+bNudcYP3587jlbtmzJPQeAA0+zZs1yz/nhD3+Ya3zt2rVzrzFnzpzccxYtWpR7DgcXdzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJIrqu4NwIFq9OjRued06dIl1/i5c+fmXuPFF1/MPQeAA0/t2rVzz9mf60bbtm1zjV+5cmXuNcaOHZt7Doc+dzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHKFLMuyfRpYKFT1XqDKDBw4MPecmTNn5p6zcePGXOP79++fe41FixblnsPBbx9fqmsc1yYOZscdd1zuOf/zP/9TBTup7Jxzzsk9Z9asWVWwEw50e7s2uaMBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJIrqu4NwP5o1qxZrvE//OEPc69Ru3bt3HPmzJmTa/yiRYtyrwHAgemYY47JNX7evHlVtJPKRo8enWv87Nmzq2gn1DTuaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJFdU3RuA2rVr554zd+7cXOPbtm2be42VK1fmnjN27NjccwA4NFx55ZW5xrdu3bqKdlLZCy+8kGt8lmVVtBNqGnc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByRdW9AWjXrl3uOSeeeGIV7KSy66+/PveclStXVsFOAPii9erVK/ec6667rgp2AgcvdzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJIrqu4NcOg55phjco2fN29eFe3k/4wePTr3nNmzZ1fBTgA4GJxyyim555SUlFTBTipbuXJl7jkfffRRFewE9s4dDQAAIDmhAQAAJCc0AACA5IQGAACQnNAAAACSExoAAEByQgMAAEhOaAAAAMkJDQAAIDmhAQAAJCc0AACA5IQGAACQXFF1b4BDz5VXXplrfOvWratoJ//nhRdeyD0ny7Iq2AkA7LBkyZLcc772ta/lnlNeXp57DqTgjgYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASK6QZVm2TwMLhareCwegXr165Z4zZ86cXONLSkpyr5HXySefnHvO4sWLq2AnsH/28aW6xnFtAqg+e7s2uaMBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJBcUXVvgAPbKaeckntOSUlJFeykspUrV+Ya/9FHH1XRTgAA2B13NAAAgOSEBgAAkJzQAAAAkhMaAABAckIDAABITmgAAADJCQ0AACA5oQEAACQnNAAAgOSEBgAAkJzQAAAAkhMaAABAckXVvQFYsmRJ7jlf+9rXco0vLy/PvQYAAPvPHQ0AACA5oQEAACQnNAAAgOSEBgAAkJzQAAAAkhMaAABAckIDAABITmgAAADJCQ0AACA5oQEAACQnNAAAgOQKWZZl+zSwUKjqvQCwB/v4Ul3juDYBVJ+9XZvc0QAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyRWyLMuqexMAAMChxR0NAAAgOaEBAAAkJzQAAIDkhAYAAJCc0AAAAJITGgAAQHJCAwAASE5oAAAAyQkNAAAguf8HThlXy+TtzGsAAAAASUVORK5CYII=",
"text/plain": [
"<Figure size 1000x500 with 2 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"torch.manual_seed(42)\n",
"\n",
"shift_x = 0.1\n",
"shift_y = 0.5\n",
"shift_transform = transforms.Compose([\n",
" transforms.RandomAffine(degrees=0, translate=(shift_x, shift_y), fill=-MNIST_MEAN/MNIST_STD, )\n",
"])\n",
"\n",
"index = 0\n",
"image, label = test_dataset[index]\n",
"\n",
"shifted_image = shift_transform(image)\n",
"\n",
"fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 5))\n",
"\n",
"show_image(image, ax1, \"Original Image\")\n",
"show_image(shifted_image, ax2, \"Shifted Image\")\n",
"\n",
"\n",
"print(f\"{'':<6}original shifted\")\n",
"print(f\"LeNet {TrainTest(leNet_clf, device).predict(image)[0]:<8} {TrainTest(leNet_clf, device).predict(shifted_image)[0]}\", )\n",
"print(f\"MLP {TrainTest(mlp_clf, device).predict(image)[0]:<8} {TrainTest(mlp_clf, device).predict(shifted_image)[0]}\")\n",
"\n",
"plt.savefig(\"../assets/Q3.png\")\n",
"plt.show()"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "base",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.11.5"
}
},
"nbformat": 4,
"nbformat_minor": 2
}

Xet Storage Details

Size:
24 kB
·
Xet hash:
4dff30408a6094068d1c68663f07a84905f7a5d606f092a30903db7c5ee37ff5

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.