{ "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.12.13", "mimetype": "text/x-python", "codemirror_mode": { "name": "ipython", "version": 3 }, "pygments_lexer": "ipython3", "nbconvert_exporter": "python", "file_extension": ".py" } }, "nbformat_minor": 4, "nbformat": 4, "cells": [ { "cell_type": "code", "source": "import matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd", "metadata": { "trusted": true, "execution": { "iopub.status.busy": "2026-07-04T16:16:09.940312Z", "iopub.execute_input": "2026-07-04T16:16:09.940612Z", "iopub.status.idle": "2026-07-04T16:16:12.856395Z", "shell.execute_reply.started": "2026-07-04T16:16:09.940586Z", "shell.execute_reply": "2026-07-04T16:16:12.855776Z" }, "ExecuteTime": { "end_time": "2026-07-04T17:01:39.086257512Z", "start_time": "2026-07-04T17:01:39.026601453Z" } }, "outputs": [], "execution_count": 46 }, { "cell_type": "code", "source": [ "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n", "#train_df=pd.read_csv(\"/kaggle/input/competitions/digit-recognizer/train.csv\")\n", "#test_df=pd.read_csv(\"/kaggle/input/competitions/digit-recognizer/test.csv\")\n", "train_df=pd.read_csv(\"train.csv\")\n", "test_df=pd.read_csv(\"test.csv\")" ], "metadata": { "trusted": true, "execution": { "iopub.status.busy": "2026-07-04T16:16:12.857730Z", "iopub.execute_input": "2026-07-04T16:16:12.858065Z", "iopub.status.idle": "2026-07-04T16:16:16.048768Z", "shell.execute_reply.started": "2026-07-04T16:16:12.858040Z", "shell.execute_reply": "2026-07-04T16:16:16.048163Z" }, "ExecuteTime": { "end_time": "2026-07-04T17:01:40.501664289Z", "start_time": "2026-07-04T17:01:39.087525035Z" } }, "outputs": [], "execution_count": 47 }, { "cell_type": "code", "source": "fig, ax = plt.subplots(nrows=2, ncols=2, sharex='all', sharey='all')\nax = ax.flatten()\nfor i in range(4):\n img = train_df.iloc[i][1:].to_numpy().reshape(28,28)\n # ax[i].imshow(img,cmap='Greys')\n ax[i].imshow(img)\n ax[i].set_title(f'{train_df.iloc[i][0]}')", "metadata": { "trusted": true, "execution": { "iopub.status.busy": "2026-07-04T16:16:16.049873Z", "iopub.execute_input": "2026-07-04T16:16:16.050089Z", "iopub.status.idle": "2026-07-04T16:16:16.369380Z", "shell.execute_reply.started": "2026-07-04T16:16:16.050066Z", "shell.execute_reply": "2026-07-04T16:16:16.368725Z" }, "ExecuteTime": { "end_time": "2026-07-04T17:01:40.828285632Z", "start_time": "2026-07-04T17:01:40.551518215Z" } }, "outputs": [ { "data": { "text/plain": [ "
" ], "image/png": "iVBORw0KGgoAAAANSUhEUgAAAeUAAAGzCAYAAAAR5w+IAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjcuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8pXeV/AAAACXBIWXMAAA9hAAAPYQGoP6dpAAAvAUlEQVR4nO3df3xU9Z3v8fdAmAHjhFSQBIKLSIDKpWCNirT8iIZ2QbGRqlTxXgjrXq/ir1i3alxbULlkxVbYDVHLrgL6qFYvGKW9EpAFuaCQIC0sUAv+CFyYJGMgbCZuSAbCuX94nZrlG2HCTOabM6/n4/F9PJzPnDnnMxi+b87ke854JDkCAAAJ1y3RDQAAgC8RygAAWIJQBgDAEoQyAACWIJQBALAEoQwAgCUIZQAALEEoAwBgCUIZAABLEMoAAFiCUHap1NRUzZs3T2vWrNHRo0flOI5mzZqV6LYAWM7r9eof/uEfFAgE1NTUpG3btmnSpEmJbitpEMou1bdvX82dO1eXXnqpdu3aleh2AHQRy5cv109/+lP95je/0QMPPKDW1la98847+v73v5/o1pKGw3Df8Hq9TkZGhiPJycnJcRzHcWbNmpXwvhgMhr3jyiuvdBzHcR566KFIzefzOR9//LHz/vvvJ7y/ZBicKbtUOBxWMBhMdBsAupCbb75ZJ0+e1NKlSyO1lpYWvfjii/re976ngQMHJrC75EAoAwAkSd/97ne1f/9+NTY2tqlXVlZKki677LIEdJVcCGUAgCSpf//+qqmpOa3+VW3AgAGd3VLSIZQBAJKkXr16qaWl5bR6c3Nz5HnEF6EMAJAkHT9+XD6f77R6z549I88jvghlAICkLz+m7t+//2n1r2rV1dWd3VLSIZQBAJKknTt3atiwYfL7/W3qY8aMiTyP+CKUAQCSpJUrVyolJUV33nlnpOb1ejV79mxt27ZNhw8fTmB3ySEl0Q0gfu655x6lp6dHVkzecMMNkesMS0pKFAqFEtkeAMtUVlbqjTfeUHFxsfr166dPPvlEs2bN0sUXX6w77rgj0e0ljYTfwYQRn1FVVeW0Z9CgQQnvj8Fg2Dd8Pp+zcOFCp7q62jl+/LhTUVHh/PCHP0x4X8kyPP//PwAAQILxO2UAACxBKAMAYAlCGQAASxDKAABYglAGAMAShDIAAJaI281D5syZo5/97GfKzMzUrl27dN9992n79u1n9doBAwac9n2egC38fj/3ALYUcwdsdjZzR1yuU54+fbpefvll3XXXXaqoqFBhYaFuueUWDR8+XHV1dd/42gEDBigQCMS6JSCmsrKyCGbLMHegKzjT3BGXUN62bZu2b9+u++6778uDeDw6dOiQSkpK9PTTT3/ja/1+v0KhkG4deKeONzbHujXgnPTy99RvDy9VWloaZ2SWYe6Azc527oj5x9c9evRQTk6OiouLIzXHcbR+/XqNHTv2tO29Xm+b7+/86ttJjjc2q6mR7+4EYMbcATeK+UKvvn37KiUlRcFgsE09GAwqMzPztO2LiooUCoUig4+fAJwN5g64UcJXXxcXFystLS0ysrKyEt0SgC6AuQNuFPOPr48cOaKTJ08qIyOjTT0jI0O1tbWnbR8OhxUOh2PdBgCXY+6AG8X8TPnEiRPasWOH8vLyIjWPx6O8vDxt3bo11ocDAMA14nKd8rPPPqsVK1boww8/VGVlpQoLC5Wamqply5bF43AAALhCXEL5jTfe0IUXXqgnn3xSmZmZ2rlzpyZPnqzPP/88HocDAMAV4nZHr9LSUpWWlsZr9wAAuE7CV18DAIAvEcoAAFiCUAYAwBKEMgAAliCUAQCwBKEMAIAlCGUAACxBKAMAYAlCGQAASxDKAABYglAGAMAShDIAAJYglAEAsAShDACAJeL21Y1Ae771/gXG+m8HbzDWRz89x1jP/McPYtYT0BWl9M801p10v7H+0QPfimr/uZd9ZKz/27KRxro35Bjrab+tMB/AMW+fzDhTBgDAEoQyAACWIJQBALAEoQwAgCViHspz586V4zhtxkcfmRcLAACAv4jL6us9e/Zo0qRJkccnT56Mx2FguYytacb6cxe9Y6yfcHoY6x4WaCJJdO9jvjIhePNwY33TzxcZ67083pj1ZLK88GNj/YbzPzXWvzf1XmN92IImY711776ONeYCcQnlkydPKhgMxmPXAAC4VlxCeejQoQoEAmpubtbWrVtVVFSkQ4cOGbf1er3y+XyRx36/+fo6APg65g64Ucx/p1xRUaGCggJNnjxZd999twYPHqzNmzfr/PPPN25fVFSkUCgUGYFAINYtAXAh5g64UcxDuby8XCtXrtTu3bu1bt06XXfddUpPT9f06dON2xcXFystLS0ysrKyYt0SABdi7oAbxf02mw0NDdq/f7+ys7ONz4fDYYXD4Xi3AcBlmDvgRnEP5dTUVA0ZMkSvvPJKvA+FBPls4Vhj/bcDf2Ws+zw+Y/3qP9xmrA9YvsdYbz2L3gAbdc/oZ6y3vmpeNV357dJ29hTfVdbtKUirbueZXsbqvtwXjfX3x5o/rJ13198a6z3/eNBYb62ra6efrifmH18/88wzmjBhggYNGqSxY8eqrKxMra2teu2112J9KAAAXCXmZ8oDBw7Ua6+9pj59+qiurk5btmzR1VdfrSNHjsT6UAAAuErMQ/m228wfQQIAgG/Gva8BALAEoQwAgCXivvoa7lE/27zKeuttvzTWz+/W01h/5ugIYz2jwLzuoDUUOovugK7j2LWXGOtbvv1cJ3eSWN/3nTLW31221FgftcR8D+2Bxay+BgAAMUYoAwBgCUIZAABLEMoAAFiCUAYAwBKsvsZpug83f3lI/oMbjfXe7ayy/rew+e7Ub//yWmM9/ejWs+gO6Dqab7jKWB96/586uZNv9p0XzKuaz6txjPVxd2831n+VWRmznkzW3L3QWJ929GfGet+lXW9O4UwZAABLEMoAAFiCUAYAwBKEMgAAliCUAQCwBKuvk9iJH15hrF/7q03G+k8v+HNU+//vCx8w1i98ueutiAQ64uQ95vu5L/ur92Ky/8c+v9xY/1//Zq63J3tDk7HueX+nsb5/ZW9j/YaM6cb6pa9+ZqwvzPzwzM19TVb384x177TPzS8w30LbapwpAwBgCUIZAABLEMoAAFiCUAYAwBJRh/L48eO1evVqBQIBOY6j/Pz807Z54oknVF1draamJr377rvKzjbfthEAAPxF1KuvU1NTtWvXLr300ksqKys77fmHH35Y999/v2bNmqWqqio99dRTWrt2rUaMGKGWlpaYNI3oBO//nrG+45Elxvopme93u/9E2Fi/40//zVjvX2ZecXnSWAW6MI/HWO7uMf9ditYV/9N8b+rUz833lx+6siImx21P6783mJ9op/7W/7naWF8w3dxnirpH1c9P/mqHsf7af5tirKe/Yu8VIFGHcnl5ucrLy9t9vrCwUPPnz9fq1aslSTNnzlQwGNSNN96o119/veOdAgDgcjH9nfLgwYPVv39/rV+/PlILhUKqqKjQ2LFjja/xer3y+/1tBgCcCXMH3CimoZyZmSlJCgaDberBYDDy3H9WVFSkUCgUGYFAIJYtAXAp5g64UcJXXxcXFystLS0ysrKyEt0SgC6AuQNuFNPbbNbW1kqSMjIyIv/91eOdO3caXxMOhxUOmxcQAUB7mDvgRjEN5aqqKtXU1CgvL0+7du2SJPn9fo0ZM0bPP/98LA8Fg5SL/8pYv/3OtTHZ/y0f/ndj/aKb9xjrrLJGsjg17jJjfePIF2Oy//7/ar63c+u+T2Ky/3jLfnCbsf79vfcb6xVPlEa1//vSzVd6lE45bqynvxLV7jtVhy6J+vp1x4MHD9bo0aNVX1+vQ4cOafHixXr88cf18ccfRy6Jqq6u1ltvvRXLvgEAcJ2oQ/mKK67Qe++9F3m8aNEiSdLy5cs1e/ZsLVy4UKmpqVq6dKnS09O1ZcsWTZ48mWuUAQA4g6hDedOmTfK0c6H8V+bOnau5c+d2uCkAAJJRwldfAwCALxHKAABYIqarr9E5umf0M9Yn/O4jY73wW/vb2ZP51xBVJ5uN9dR3uGMSYPLv2T1jsp9PT5pXC3vCJ2Kyf9tkbKgx1j/9ufnPYUhKr3i2YwXOlAEAsAShDACAJQhlAAAsQSgDAGAJQhkAAEuw+rorSjvfWP7pBX+Oye4LL7/BWL/g6NaY7B9wm57/fiom+3ns/+Yb66eCdTHZv21OfnbAWL91198Y69tzXotq/89cudJYX/qtK4311mPHotp/PHCmDACAJQhlAAAsQSgDAGAJQhkAAEsQygAAWILV1xZLGZhlrF+10rzKuls797Juz4M1Y4x157j53tdAsuvet4+x/g+/ej4m+3/9knXG+g0XTTe/YN8nMTmubbxvfMv8RE50+7nhvJCx/s8+b5QddR7OlAEAsAShDACAJQhlAAAsQSgDAGCJqEN5/PjxWr16tQKBgBzHUX5+29vCLVu2TI7jtBlr1qyJWcMAALhV1KuvU1NTtWvXLr300ksqKyszbrNmzRrNnj078rilpaXjHSaxz19INdYf67vbWG/v7rsPVH/fWK+aaP432ammpjP2BiQjT48exvrVvk5uxOX8h5I3M6IO5fLycpWXl3/jNi0tLQoGgx1uCgCAZBSX65Rzc3MVDAZ17NgxbdiwQY8//rjq6+uN23q9Xvl8f/lnpt/vj0dLAFyGuQNuFPOFXuXl5Zo5c6by8vL0yCOPaOLEiVqzZo26dTMfqqioSKFQKDICgUCsWwLgQswdcKOYh/Lrr7+u3/3ud9qzZ4/efvttTZ06VVdddZVyc3ON2xcXFystLS0ysrLMd7ECgK9j7oAbxf02m1VVVaqrq1N2drY2bNhw2vPhcFjhcDjebQBwGeYOuFHcQzkrK0t9+vRRTU1NvA/VZbV3j+sfZJnvcd2eL06ZVyzu+KfvGuvpTVuj2j+Q7E4G64z1726/3Vj/45W/iWc7cKEOXRKVnZ0deTx48GCNHj1a9fX1qq+v19y5c7Vq1SrV1tZqyJAhWrhwoT755BOtXbs2po0DAOA2UYfyFVdcoffeey/yeNGiRZKk5cuX6+6779aoUaM0a9Yspaenq7q6WuvWrdPPf/5zPmYCAOAMog7lTZs2yeNp/ysCJ0+efE4NAQCQrLj3NQAAliCUAQCwRNxXX+MvUgZdZKz7X/0PY/2Jfn801o+0HjfWp/zyYWM945UPzqI7AGd0qtVY9mz8lnn7K2Nz2Etf/cxY/2iS+bitx47F5sBx1j2jn7F+7ZItMdn/sI13GOvZwZ0x2X88cKYMAIAlCGUAACxBKAMAYAlCGQAASxDKAABYgtXXnejgbebV13+8uCSq/TwSuM5Yz/gnVlkDiZD16sfG+vy/GWmsP953T1T7X5j5obH+2IbLjfX3548x1lNXVUR13FhJuWigsX7wH3sb6393QXlU+/+8tclYH77AfGVLq+NEtf/OxJkyAACWIJQBALAEoQwAgCUIZQAALEEoAwBgCVZfx8Hnc75nrL959zPtvKKnsXpvYJyxfvT2C9rZT+gMnQGIh9a6OmN9w9+b/w73ftq8Wvi+dPM9rtuzoN8fjPW7Hk411g8c+W5U+085Zr7P/qmePcz1XuZImdDOvaz/7oJ9UfXTnh/vnWWsp/1pf0z235k4UwYAwBKEMgAAliCUAQCwBKEMAIAlogrlRx99VJWVlQqFQgoGgyorK9OwYcPabOPz+bRkyRIdOXJEjY2NWrlypfr1M3+RNQAA+IuoVl9PnDhRpaWl2r59u1JSUrRgwQKtW7dOI0aMUFPTl6sJFy1apOuvv1633HKLGhoatGTJEr355psaN868CrEr637hhcb63z3wurE+OMW8yro9f3j+MmP9gs+2RrUfAInR8/eVxvorWVOM9R//vfkKjazu50V13BcGbjY/8Wo79XZsbzHfI3pAinlVdrR9xkr4rfZO/D7t1D5iIapQnjKl7Q9SQUGB6urqlJOTo82bNystLU133HGHZsyYoY0bN0qSZs+erT//+c8aM2aMKioSczN0AAC6gnO6Trl37y+/4aO+vl6SlJOTI6/Xq/Xr10e22bdvnw4ePKixY8caQ9nr9crn80Ue+/3+c2kJQJJg7oAbdXihl8fj0eLFi7Vlyxbt3btXkpSZmamWlhY1NDS02TYYDCozM9O4n6KiIoVCocgIBAIdbQlAEmHugBt1OJRLS0s1cuRI3XrrrefUQHFxsdLS0iIjKyvrnPYHIDkwd8CNOvTxdUlJiaZOnaoJEya0+ddpbW2tfD6fevfu3eZsOSMjQ7W1tcZ9hcNhhcPhjrQBIIkxd8CNog7lkpISTZs2Tbm5uTpw4ECb53bs2KFwOKy8vDy9+eabkqRhw4Zp0KBB2rrVfSuGAzOGGuvTzy+Pyf7DaZ6Y7AeAXfr+2jwf/jDrZ8b63jtK49lOu670tTcHxXeV9f4Tzcb6fy1+yFjPeP1PxnprzDrqPFGFcmlpqWbMmKH8/Hw1NjYqIyNDktTQ0KDm5maFQiG9+OKLevbZZ1VfX69QKKSSkhJ98MEHrLwGAOAMogrlOXPmSJI2bdrUpl5QUKAVK1ZIkh588EGdOnVKq1atks/n09q1ayOvAwAA7YsqlD2eM3+c2tLSonvvvVf33ntvh5sCACAZce9rAAAsQSgDAGCJc7qjV7LrdsJcP+GY1/z18HQ31lsc844ah5j3Y74NC4Cu7pJ/3Ges50+43lh/e+j/jmc7cRdobTLW73jk74z1vq+bV613xVXW7eFMGQAASxDKAABYglAGAMAShDIAAJYglAEAsASrr89Bv+c+MNaX3TvEWE/t1mKsL3rhZmN96GLz/gG4U+vRemPduT7VWP/ej+8x1uvyzF/U8fEP/tlY7+4xn5+1Oqei2v6SdXcY65f+fY2x7oTNV57467YZ68mAM2UAACxBKAMAYAlCGQAASxDKAABYglAGAMASrL6Og9Uj+kS1faZYZQ2gfaf+4z+M9fRXzPeCTn/FvJ/rdHmsWjIaqh3G+sm4HtVdOFMGAMAShDIAAJYglAEAsAShDACAJaIK5UcffVSVlZUKhUIKBoMqKyvTsGHD2myzceNGOY7TZjz//PMxbRoAADeKKpQnTpyo0tJSXX311frBD36gHj16aN26dTrvvPPabLd06VJlZmZGxsMPPxzTpgEAcKOoLomaMmVKm8cFBQWqq6tTTk6ONm/eHKk3NTUpGAzGpkMAAJLEOf1OuXfv3pKk+vq232xy++23q66uTrt379aCBQvUq1evdvfh9Xrl9/vbDAA4E+YOuFGHbx7i8Xi0ePFibdmyRXv37o3UX331VR08eFDV1dUaNWqUnn76aQ0fPlw33XSTcT9FRUWaN29eR9sAkKSYO+BGHklOR1743HPPacqUKRo3bpwCgUC7211zzTXasGGDhgwZos8+++y0571er3w+X+Sx3+9XIBBQfu+Zamo83pHWgLg5z99Lbze8rLS0NDU2Nia6naTG3IGu5Gznjg6dKZeUlGjq1KmaMGHCNwayJFVUVEiSsrOzjaEcDocVDpu/kBsA2sPcATeKOpRLSko0bdo05ebm6sCBA2fc/rLLLpMk1dTURHsoAACSSlShXFpaqhkzZig/P1+NjY3KyMiQJDU0NKi5uVmXXHKJZsyYoXfeeUdHjx7VqFGjtGjRIm3atEm7d++OyxsAAMAtogrlOXPmSJI2bdrUpl5QUKAVK1YoHA5r0qRJKiwsVGpqqg4dOqRVq1Zp/vz5sesYAACXiiqUPR7PNz5/+PBh5ebmnks/AAAkLe59DQCAJQhlAAAsQSgDAGAJQhkAAEsQygAAWIJQBgDAEoQyAACW6PC3RMVbL3/PRLcAnIafS/vx/wg2Otufyw5/S1S8DBgw4IxfcgEkWlZWlqqrqxPdBr6GuQNdwZnmDutCWfryL1djY2Pkq9iysrKS4mvyeL9dg9/vJ5AtxdzB+7XZ2cwdVn58/Z+bbmxs7FJ/8OeK92u3rtRrsmHu4P3a7Gx6ZaEXAACWIJQBALCE1aHc0tKiefPmqaWlJdGtdAreLxAbyfazxft1DysXegEAkIysPlMGACCZEMoAAFiCUAYAwBKEMgAAliCUAQCwBKEMAIAlCGWXSk1N1bx587RmzRodPXpUjuNo1qxZiW4LQBfy2GOPyXEc7d69O9GtJA1C2aX69u2ruXPn6tJLL9WuXbsS3Q6ALiYrK0uPPfaYvvjii0S3klSs/EIKnLuamhplZmYqGAwqJydHH374YaJbAtCF/PKXv9S2bdvUvXt39e3bN9HtJA3OlF0qHA4rGAwmug0AXdD48eN18803q7CwMNGtJB1CGQAQ0a1bN5WUlOhf/uVftGfPnkS3k3T4+BoAEHHXXXdp0KBBmjRpUqJbSUqcKQMAJEkXXHCBnnzyST311FM6cuRIottJSoQyAECSNH/+fNXX16ukpCTRrSQtPr4GACg7O1t33nmnCgsLNWDAgEi9Z8+e6tGjhwYNGqRQKKRjx44lsEv340wZAKCsrCx1795dJSUlOnDgQGRcffXVGj58uA4cOKBf/OIXiW7T9ThTBgBoz549uvHGG0+rz58/X36/Xw888IA+/fTTzm8syXgkOYluAvFxzz33KD09XQMGDNCcOXO0atUq/fGPf5QklZSUKBQKJbhDALbbuHGj+vbtq+985zuJbiUpEMouVlVVpYsvvtj43MUXX6yDBw92bkMAuhxCuXMRygAAWIKFXgAAWIJQBgDAEoQyAACWIJQBALAEoQwAgCXidvOQOXPm6Gc/+5kyMzO1a9cu3Xfffdq+fftZvXbAgAFqbGyMV2vAOfH7/aqurk50GzBg7oDNzmbuiMslUdOnT9fLL7+su+66SxUVFSosLNQtt9yi4cOHq66u7htfO2DAAAUCgVi3BMRUVlYWwWwZ5g50BWeaO+ISytu2bdP27dt13333fXkQj0eHDh1SSUmJnn766W98rd/vVygU0q0D79TxxuZYtwack17+nvrt4aVKS0vjjMwyzB2w2dnOHTH/+LpHjx7KyclRcXFxpOY4jtavX6+xY8eetr3X65XP54s89vv9kqTjjc1qajwe6/YAuARzB9wo5gu9+vbtq5SUFAWDwTb1YDCozMzM07YvKipSKBSKDD5+AnA2mDvgRglffV1cXKy0tLTIyMrKSnRLALoA5g64Ucw/vj5y5IhOnjypjIyMNvWMjAzV1taetn04HFY4HI51GwBcjrkDbhTzM+UTJ05ox44dysvLi9Q8Ho/y8vK0devWWB8OAADXiMt1ys8++6xWrFihDz/8UJWVlSosLFRqaqqWLVsWj8MBAOAKcQnlN954QxdeeKGefPJJZWZmaufOnZo8ebI+//zzeBwOAABXiNsdvUpLS1VaWhqv3QMA4DoJX30NAAC+RCgDAGAJQhkAAEsQygAAWIJQBgDAEoQyAACWIJQBALAEoQwAgCUIZQAALEEoAwBgCUIZAABLEMoAAFiCUAYAwBKEMgAAlojbVzcifjwp5v9t+57/rvkFp9rf1/B7/misOydPRtsWAOAccaYMAIAlCGUAACxBKAMAYAlCGQAAS8Q8lOfOnSvHcdqMjz76KNaHAQDAdeKy+nrPnj2aNGlS5PFJVvLGlKdXL2P9k+t+HfW+pv50vLHO6mvg3M3cd8hYf/nwWGO92/VHjPVTzc0x6ykRuvn9xnr9tJHGevrLW+PZjtXiEsonT55UMBiMx64BAHCtuITy0KFDFQgE1NzcrK1bt6qoqEiHDpn/xej1euXz+SKP/e38iwoAvo65A24U898pV1RUqKCgQJMnT9bdd9+twYMHa/PmzTr//PON2xcVFSkUCkVGIBCIdUsAXIi5A24U81AuLy/XypUrtXv3bq1bt07XXXed0tPTNX36dOP2xcXFSktLi4ysrKxYtwTAhZg74EZxv81mQ0OD9u/fr+zsbOPz4XBY4XA43m0AcBnmDrhR3EM5NTVVQ4YM0SuvvBLvQwGAVX5zQ665/q8vG+uz0n9srJ+q7dqrrz2ZFxrruQ+aV1nvNP/xJIWYf3z9zDPPaMKECRo0aJDGjh2rsrIytba26rXXXov1oQAAcJWYnykPHDhQr732mvr06aO6ujpt2bJFV199tY4cMV9/BwAAvhTzUL7ttttivUsAAJIC974GAMAShDIAAJaI++pr2O3/PjDaWB+44INO7gRwn9b9nxrrjaccY/3jxRnG+uBb3Xnb4gX9/mCsX3PjXcZ6r7cq49mOFThTBgDAEoQyAACWIJQBALAEoQwAgCUIZQAALMHq6yQ39K/Nq0OPL+jkRoAkMvXD/2GszxxhXl38fs90Y/1Uc9e+J3Z7nG6eRLeQMJwpAwBgCUIZAABLEMoAAFiCUAYAwBKEMgAAlmD1NQB0suaDfmO96Oo/Ges/uvBHxvqpQ4dj1lM8eY63GOv7T7hz9fi54EwZAABLEMoAAFiCUAYAwBKEMgAAlog6lMePH6/Vq1crEAjIcRzl5+efts0TTzyh6upqNTU16d1331V2dnZMmgUAwM2iXn2dmpqqXbt26aWXXlJZWdlpzz/88MO6//77NWvWLFVVVempp57S2rVrNWLECLW0mFfgAUAy6buznXs7/6Rz++gsJw8HjPXFn+d1cif2izqUy8vLVV5e3u7zhYWFmj9/vlavXi1JmjlzpoLBoG688Ua9/vrrHe8UAACXi+nvlAcPHqz+/ftr/fr1kVooFFJFRYXGjh1rfI3X65Xf728zAOBMmDvgRjEN5czMTElSMBhsUw8Gg5Hn/rOioiKFQqHICATMH3MAwNcxd8CNEr76uri4WGlpaZGRlZWV6JYAdAHMHXCjmN5ms7a2VpKUkZER+e+vHu/cudP4mnA4rHA4HMs2ACQB5g64UUxDuaqqSjU1NcrLy9OuXbskSX6/X2PGjNHzzz8fy0MltxMnjOVbPv1rY/1/DVkbz24ARKl7i5PoFqx2+LpWY33Ym53cSAJ06JKor193PHjwYI0ePVr19fU6dOiQFi9erMcff1wff/xx5JKo6upqvfXWW7HsGwAA14k6lK+44gq99957kceLFi2SJC1fvlyzZ8/WwoULlZqaqqVLlyo9PV1btmzR5MmTuUYZAIAziDqUN23aJI+nnQvf/7+5c+dq7ty5HW4KAIBklPDV1wAA4EuEMgAAlojp6mt0jlPNzcZ61W8vN7/g71l9DdjE12BeXdzinOzkTuz0fO4rxvoiXdrJnXQ+zpQBALAEoQwAgCUIZQAALEEoAwBgCUIZAABLsPq6C/L08BrrDVdx1zSgK/CWbzfWf990obG+/+m+xvqQ2XXGutNF7qC4ccNlxvpDt6031rv3ucBYbz1aH6uWEo4zZQAALEEoAwBgCUIZAABLEMoAAFiCUAYAwBKsvu6CPD19xvrHP/jnTu4EQCz902O3Guu7FpcY6z8edYd5R9t3x6qluOpVY/4a4GE9Uo31hrxhxvr5b2yLWU+JxpkyAACWIJQBALAEoQwAgCUIZQAALBF1KI8fP16rV69WIBCQ4zjKz89v8/yyZcvkOE6bsWbNmpg1DACAW0W9+jo1NVW7du3SSy+9pLKyMuM2a9as0ezZsyOPW7rIfVgBIJFSV1YY63ueMa9S7vnLz4314xNj1lJcDVx5wFiveeiLzm3EIlGHcnl5ucrLy79xm5aWFgWDwQ43BQBAMorLdcq5ubkKBoM6duyYNmzYoMcff1z19eZv8fB6vfL5/nLdrd/vj0dLAFyGuQNuFPOFXuXl5Zo5c6by8vL0yCOPaOLEiVqzZo26dTMfqqioSKFQKDICgUCsWwLgQswdcKOYh/Lrr7+u3/3ud9qzZ4/efvttTZ06VVdddZVyc3ON2xcXFystLS0ysrKyYt0SABdi7oAbxf02m1VVVaqrq1N2drY2bNhw2vPhcFjhcDjebQBwGeYOuFHcQzkrK0t9+vRRTU1NvA8FAEml+os0Y/1b6hoLbVuD5tXjT9flGuvfmnPQWD9Vbv5zaA2FOtRXInXokqjs7OzI48GDB2v06NGqr69XfX295s6dq1WrVqm2tlZDhgzRwoUL9cknn2jt2rUxbRwAALeJOpSvuOIKvffee5HHixYtkiQtX75cd999t0aNGqVZs2YpPT1d1dXVWrdunX7+85/zMRMAAGcQdShv2rRJHo/5QnZJmjx58jk1BABAsuLe1wAAWIJQBgDAEnFffQ0AODf/ddvfGuu3jfjQWK/okWqsOyeiW9vTPXuwsX7sygxj/fOrzPv5Se4Hxvr53RuN9Uf6fGTeUaa5PHT+3eb6/eZ7iduMM2UAACxBKAMAYAlCGQAASxDKAABYglAGAMASrL4GAMv1f9VnrP/ihd3G+rBn5hjrPRrM52Ejr91vrJcMesVY793Na6z/7cG/NtY3/Op7xnqvI63G+j/nTzTWP/nRC8Z6xrb2b2jV1XCmDACAJQhlAAAsQSgDAGAJQhkAAEsQygAAWILV113QZ/9ivh+ttKlT+wDQOVK3VRnrL4YGGuu/+VFpVPv/mz/MMtYnvfOwsZ5Z2WKsp/zrDmO9t7ZF1c/wuv9ifuJHUe2mS+JMGQAASxDKAABYglAGAMAShDIAAJaIKpQfffRRVVZWKhQKKRgMqqysTMOGDWuzjc/n05IlS3TkyBE1NjZq5cqV6tevX0ybBgDAjaJafT1x4kSVlpZq+/btSklJ0YIFC7Ru3TqNGDFCTU1NkqRFixbp+uuv1y233KKGhgYtWbJEb775psaNGxeXN5CM/kv/GmO9u4cPPgA3aq2rM9ZXXWo+4Vml6E6ELtKeqHuKp+7VRxPdQsJEFcpTpkxp87igoEB1dXXKycnR5s2blZaWpjvuuEMzZszQxo0bJUmzZ8/Wn//8Z40ZM0YVFRWx6xwAAJc5p+uUe/fuLUmqr6+XJOXk5Mjr9Wr9+vWRbfbt26eDBw9q7NixxlD2er3y+f7yDSh+v/9cWgKQJJg74EYd/rzT4/Fo8eLF2rJli/bu3StJyszMVEtLixoaGtpsGwwGlZmZadxPUVGRQqFQZAQCgY62BCCJMHfAjTocyqWlpRo5cqRuvfXWc2qguLhYaWlpkZGVlXVO+wOQHJg74EYd+vi6pKREU6dO1YQJE9r867S2tlY+n0+9e/duc7ackZGh2tpa477C4bDC4XBH2gCQxJg74EZRnymXlJRo2rRpuvbaa3XgwIE2z+3YsUPhcFh5eXmR2rBhwzRo0CBt3br1nJvFN2t1TkU9AAD2iOpMubS0VDNmzFB+fr4aGxuVkZEhSWpoaFBzc7NCoZBefPFFPfvss6qvr1coFFJJSYk++OADVl4DAHAGUYXynDlzJEmbNrX9NqKCggKtWLFCkvTggw/q1KlTWrVqlXw+n9auXRt5HQAAaF9UoezxeM64TUtLi+69917de++9HW4KAIBkxC2gAACwBKEMAIAlzumOXgAAxFpr/TFjff6RkcZ66GLz+WVazDrqPJwpAwBgCUIZAABLEMoAAFiCUAYAwBKEMgAAlmD1dRd05NnB5idKo99X/bODjPVeCka/MwCIAaelxVjfHRpg3v7yUDzb6VScKQMAYAlCGQAASxDKAABYglAGAMAShDIAAJZg9XUX1OutSmP9urcuj35fMu8LABKlW8+exvqV6QeN9X2/GxbPdjoVZ8oAAFiCUAYAwBKEMgAAliCUAQCwRFSh/Oijj6qyslKhUEjBYFBlZWUaNqztL9g3btwox3HajOeffz6mTQMA4EZRrb6eOHGiSktLtX37dqWkpGjBggVat26dRowYoaampsh2S5cu1S9+8YvI468/BwDANznV3Gysb/hOqrE+QB/Es51OFVUoT5kypc3jgoIC1dXVKScnR5s3b47Um5qaFAzyhQYAAETjnH6n3Lt3b0lSfX19m/rtt9+uuro67d69WwsWLFCvXr3a3YfX65Xf728zAOBMmDvgRh2+eYjH49HixYu1ZcsW7d27N1J/9dVXdfDgQVVXV2vUqFF6+umnNXz4cN10003G/RQVFWnevHkdbQNAkmLugBt5JDkdeeFzzz2nKVOmaNy4cQoEAu1ud80112jDhg0aMmSIPvvss9Oe93q98vl8kcd+v1+BQED5vWeqqfF4R1oD4uY8fy+93fCy0tLS1NjYmOh2khpzB7qSs507OnSmXFJSoqlTp2rChAnfGMiSVFFRIUnKzs42hnI4HFY4HO5IGwCSGHMH3CjqUC4pKdG0adOUm5urAwcOnHH7yy67TJJUU1MT7aEAAEgqUYVyaWmpZsyYofz8fDU2NiojI0OS1NDQoObmZl1yySWaMWOG3nnnHR09elSjRo3SokWLtGnTJu3evTsubwAAALeIKpTnzJkjSdq0aVObekFBgVasWKFwOKxJkyapsLBQqampOnTokFatWqX58+fHrmMAAFwqqlD2eDzf+Pzhw4eVm5t7Lv0AAJC0uPc1AACWIJQBALAEoQwAgCUIZQAALEEoAwBgCUIZAABLEMoAAFiiw98SFW+9/D0T3QJwGn4u7cf/I9jobH8uO/wtUfEyYMCAM37JBZBoWVlZqq6uTnQb+BrmDnQFZ5o7rAtl6cu/XI2NjZGvYsvKykqKr8nj/XYNfr+fQLYUcwfv12ZnM3dY+fH1f266sbGxS/3Bnyver926Uq/JhrmD92uzs+mVhV4AAFiCUAYAwBJWh3JLS4vmzZunlpaWRLfSKXi/QGwk288W79c9rFzoBQBAMrL6TBkAgGRCKAMAYAlCGQAASxDKAABYglAGAMASVofynDlzVFVVpePHj2vbtm268sorE91STIwfP16rV69WIBCQ4zjKz88/bZsnnnhC1dXVampq0rvvvqvs7OwEdBobjz76qCorKxUKhRQMBlVWVqZhw4a12cbn82nJkiU6cuSIGhsbtXLlSvXr1y9BHaMrc+u8ISXX3JHM84Zj45g+fbrT3NzsFBQUOJdeeqnz61//2qmvr3cuvPDChPd2rmPy5MnOU0895dx4442O4zhOfn5+m+cffvhh59ixY86PfvQj5zvf+Y7z1ltvOZ9++qnj8/kS3ntHxpo1a5xZs2Y5I0aMcEaNGuX8/ve/dw4cOOCcd955kW2ee+455+DBg84111zjXH755c4HH3zgbNmyJeG9M7rWcPO8ISXX3JHE80bCGzCObdu2OSUlJZHHHo/HOXz4sPPII48kvLdYDtNfrOrqauehhx6KPE5LS3OOHz/u/OQnP0l4v7EYffv2dRzHccaPHx95fy0tLc5NN90U2Wb48OGO4zjOmDFjEt4vo+uMZJk3pOSbO5Jl3rDy4+sePXooJydH69evj9Qcx9H69es1duzYBHYWf4MHD1b//v3bvPdQKKSKigrXvPfevXtLkurr6yVJOTk58nq9bd7zvn37dPDgQde8Z8RfMs8bkvvnjmSZN6wM5b59+yolJUXBYLBNPRgMKjMzM0FddY6v3p9b37vH49HixYu1ZcsW7d27V9KX77mlpUUNDQ1ttnXLe0bnSOZ5Q3L33JFM84aVX90I9yotLdXIkSM1bty4RLcCoItIpnnDyjPlI0eO6OTJk8rIyGhTz8jIUG1tbYK66hxfvT83vveSkhJNnTpV11xzjQKBQKReW1srn88X+XjqK254z+g8yTxvSO6dO5Jt3rAylE+cOKEdO3YoLy8vUvN4PMrLy9PWrVsT2Fn8VVVVqaamps179/v9GjNmTJd+7yUlJZo2bZquvfZaHThwoM1zO3bsUDgcbvOehw0bpkGDBnXp94zOlczzhuTOuSNZ542ErzYzjenTpzvHjx93Zs6c6Xz72992XnjhBae+vt7p169fwns715GamuqMHj3aGT16tOM4jlNYWOiMHj3aueiiixzpy8sa6uvrnRtuuMEZOXKkU1ZW1mUva5DklJaWOseOHXMmTJjgZGRkREbPnj0j2zz33HPOgQMHnNzcXOfyyy933n//fef9999PeO+MrjXcPG9IyTV3JPG8kfAG2h333HOPc+DAAae5udnZtm2bc9VVVyW8p1iMiRMnOibLli2LbPPEE084NTU1zvHjx513333XGTp0aML77uhoz6xZsyLb+Hw+Z8mSJc7Ro0edL774wlm1apWTkZGR8N4ZXW+4dd6QkmvuSNZ5g+9TBgDAElb+ThkAgGREKAMAYAlCGQAASxDKAABYglAGAMAShDIAAJYglAEAsAShDACAJQhlAAAsQSgDAGAJQhkAAEv8Pw9w43TxJfjXAAAAAElFTkSuQmCC" }, "metadata": {}, "output_type": "display_data", "jetTransient": { "display_id": null } } ], "execution_count": 48 }, { "cell_type": "code", "source": "class DatasetMnist(Dataset):\n def __init__(self,df):\n self.df=df\n def __len__(self):\n return len(self.df)\n def __getitem__(self, idx):\n item= {\"Data\": torch.tensor(self.df.iloc[idx][1:].to_numpy().reshape(28, 28),dtype=torch.float), \"label\": torch.tensor(self.df.iloc[idx][0],dtype=torch.long)}\n return item\nbatch_size =64\ntrain_dataset = DatasetMnist(train_df)\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)", "metadata": { "trusted": true, "execution": { "iopub.status.busy": "2026-07-04T16:16:16.370253Z", "iopub.execute_input": "2026-07-04T16:16:16.370426Z", "iopub.status.idle": "2026-07-04T16:16:16.375630Z", "shell.execute_reply.started": "2026-07-04T16:16:16.370409Z", "shell.execute_reply": "2026-07-04T16:16:16.375171Z" }, "ExecuteTime": { "end_time": "2026-07-04T17:01:40.864484691Z", "start_time": "2026-07-04T17:01:40.830443307Z" } }, "outputs": [], "execution_count": 49 }, { "cell_type": "code", "source": [ "class MnistModule(nn.Module):\n", " '''\n", " (batch_size,1,28,28)\n", " '''\n", " def __init__(self):\n", " super(MnistModule, self).__init__()\n", " self.cov1=nn.Conv2d(in_channels=1, out_channels=16, kernel_size=7)\n", " self.cov2=nn.Conv2d(in_channels=16, out_channels=32, kernel_size=7)\n", " self.cov3=nn.Conv2d(in_channels=32, out_channels=64, kernel_size=5)\n", " self.cov4=nn.Conv2d(in_channels=64, out_channels=128, kernel_size=5)\n", " self.cov5=nn.Conv2d(in_channels=128, out_channels=256, kernel_size=5)\n", " self.cov6=nn.Conv2d(in_channels=256, out_channels=512, kernel_size=4)\n", " self.linear1 = nn.Linear(in_features=1*1*512, out_features=128)\n", " self.linear2 = nn.Linear(in_features=128, out_features=10)\n", " self.relu = nn.ReLU()\n", "\n", " def forward(self,X):\n", " X=X.view(-1,1,28,28)\n", " X=self.cov6(self.cov5(self.cov4(self.cov3(self.cov2(self.cov1(X))))))\n", " X=X.view(-1,512)\n", " return self.linear2(self.relu(self.linear1(X)))\n" ], "metadata": { "trusted": true, "execution": { "iopub.status.busy": "2026-07-04T16:16:16.376258Z", "iopub.execute_input": "2026-07-04T16:16:16.376401Z", "iopub.status.idle": "2026-07-04T16:16:16.389616Z", "shell.execute_reply.started": "2026-07-04T16:16:16.376386Z", "shell.execute_reply": "2026-07-04T16:16:16.388703Z" }, "ExecuteTime": { "end_time": "2026-07-04T17:05:35.534234524Z", "start_time": "2026-07-04T17:05:35.497304936Z" } }, "outputs": [], "execution_count": 56 }, { "cell_type": "code", "source": [ "model = MnistModule()\n", "model=model.to(device)\n", "loss_func = nn.CrossEntropyLoss()\n", "optimizer = torch.optim.AdamW(model.parameters(), lr=2e-4)\n", "scheduler=torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)" ], "metadata": { "trusted": true, "execution": { "iopub.status.busy": "2026-07-04T16:16:16.390304Z", "iopub.execute_input": "2026-07-04T16:16:16.390504Z", "iopub.status.idle": "2026-07-04T16:16:18.580638Z", "shell.execute_reply.started": "2026-07-04T16:16:16.390482Z", "shell.execute_reply": "2026-07-04T16:16:18.580038Z" }, "ExecuteTime": { "end_time": "2026-07-04T17:05:36.050933953Z", "start_time": "2026-07-04T17:05:36.028103516Z" } }, "outputs": [], "execution_count": 57 }, { "cell_type": "code", "source": "epochs = 30\nfor epoch in range(epochs):\n model.train()\n training_loss=0\n for batch in train_loader:\n optimizer.zero_grad()\n X = batch['Data']\n labels=batch['label']\n X=X.to(device)\n labels=labels.to(device)\n #print(model(X),'\\n',labels)\n loss = loss_func(model(X),labels)\n loss.backward()\n optimizer.step()\n training_loss+=loss.item()\n scheduler.step()\n print(f\"train_loss: {training_loss/len(train_loader)}\")", "metadata": { "trusted": true, "execution": { "iopub.status.busy": "2026-07-04T16:16:18.582512Z", "iopub.execute_input": "2026-07-04T16:16:18.582874Z", "iopub.status.idle": "2026-07-04T16:18:53.082304Z", "shell.execute_reply.started": "2026-07-04T16:16:18.582856Z", "shell.execute_reply": "2026-07-04T16:18:53.081057Z" }, "ExecuteTime": { "end_time": "2026-07-04T17:05:48.843650166Z", "start_time": "2026-07-04T17:05:36.590087546Z" } }, "outputs": [ { "ename": "KeyboardInterrupt", "evalue": "", "output_type": "error", "traceback": [ "\u001B[31m---------------------------------------------------------------------------\u001B[39m", "\u001B[31mKeyboardInterrupt\u001B[39m Traceback (most recent call last)", "\u001B[36mCell\u001B[39m\u001B[36m \u001B[39m\u001B[32mIn[58]\u001B[39m\u001B[32m, line 14\u001B[39m\n\u001B[32m 12\u001B[39m loss = loss_func(model(X),labels)\n\u001B[32m 13\u001B[39m loss.backward()\n\u001B[32m---> \u001B[39m\u001B[32m14\u001B[39m \u001B[43moptimizer\u001B[49m\u001B[43m.\u001B[49m\u001B[43mstep\u001B[49m\u001B[43m(\u001B[49m\u001B[43m)\u001B[49m\n\u001B[32m 15\u001B[39m training_loss+=loss.item()\n\u001B[32m 16\u001B[39m scheduler.step()\n", "\u001B[36mFile \u001B[39m\u001B[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/optim/lr_scheduler.py:166\u001B[39m, in \u001B[36mLRScheduler.__init__..patch_track_step_called..wrap_step..wrapper\u001B[39m\u001B[34m(*args, **kwargs)\u001B[39m\n\u001B[32m 164\u001B[39m opt = opt_ref()\n\u001B[32m 165\u001B[39m opt._opt_called = \u001B[38;5;28;01mTrue\u001B[39;00m \u001B[38;5;66;03m# type: ignore[union-attr]\u001B[39;00m\n\u001B[32m--> \u001B[39m\u001B[32m166\u001B[39m \u001B[38;5;28;01mreturn\u001B[39;00m \u001B[43mfunc\u001B[49m\u001B[43m.\u001B[49m\u001B[34;43m__get__\u001B[39;49m\u001B[43m(\u001B[49m\u001B[43mopt\u001B[49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[43mopt\u001B[49m\u001B[43m.\u001B[49m\u001B[34;43m__class__\u001B[39;49m\u001B[43m)\u001B[49m\u001B[43m(\u001B[49m\u001B[43m*\u001B[49m\u001B[43margs\u001B[49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[43m*\u001B[49m\u001B[43m*\u001B[49m\u001B[43mkwargs\u001B[49m\u001B[43m)\u001B[49m\n", "\u001B[36mFile \u001B[39m\u001B[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/optim/optimizer.py:526\u001B[39m, in \u001B[36mOptimizer.profile_hook_step..wrapper\u001B[39m\u001B[34m(*args, **kwargs)\u001B[39m\n\u001B[32m 521\u001B[39m \u001B[38;5;28;01mraise\u001B[39;00m \u001B[38;5;167;01mRuntimeError\u001B[39;00m(\n\u001B[32m 522\u001B[39m \u001B[33mf\u001B[39m\u001B[33m\"\u001B[39m\u001B[38;5;132;01m{\u001B[39;00mfunc\u001B[38;5;132;01m}\u001B[39;00m\u001B[33m must return None or a tuple of (new_args, new_kwargs), but got \u001B[39m\u001B[38;5;132;01m{\u001B[39;00mresult\u001B[38;5;132;01m}\u001B[39;00m\u001B[33m.\u001B[39m\u001B[33m\"\u001B[39m\n\u001B[32m 523\u001B[39m )\n\u001B[32m 525\u001B[39m \u001B[38;5;66;03m# pyrefly: ignore [invalid-param-spec]\u001B[39;00m\n\u001B[32m--> \u001B[39m\u001B[32m526\u001B[39m out = \u001B[43mfunc\u001B[49m\u001B[43m(\u001B[49m\u001B[43m*\u001B[49m\u001B[43margs\u001B[49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[43m*\u001B[49m\u001B[43m*\u001B[49m\u001B[43mkwargs\u001B[49m\u001B[43m)\u001B[49m\n\u001B[32m 527\u001B[39m \u001B[38;5;28mself\u001B[39m._optimizer_step_code()\n\u001B[32m 529\u001B[39m \u001B[38;5;66;03m# call optimizer step post hooks\u001B[39;00m\n", "\u001B[36mFile \u001B[39m\u001B[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/optim/optimizer.py:81\u001B[39m, in \u001B[36m_use_grad_for_differentiable.._use_grad\u001B[39m\u001B[34m(*args, **kwargs)\u001B[39m\n\u001B[32m 79\u001B[39m torch.set_grad_enabled(\u001B[38;5;28mself\u001B[39m.defaults[\u001B[33m\"\u001B[39m\u001B[33mdifferentiable\u001B[39m\u001B[33m\"\u001B[39m])\n\u001B[32m 80\u001B[39m torch._dynamo.graph_break()\n\u001B[32m---> \u001B[39m\u001B[32m81\u001B[39m ret = \u001B[43mfunc\u001B[49m\u001B[43m(\u001B[49m\u001B[43m*\u001B[49m\u001B[43margs\u001B[49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[43m*\u001B[49m\u001B[43m*\u001B[49m\u001B[43mkwargs\u001B[49m\u001B[43m)\u001B[49m\n\u001B[32m 82\u001B[39m \u001B[38;5;28;01mfinally\u001B[39;00m:\n\u001B[32m 83\u001B[39m torch._dynamo.graph_break()\n", "\u001B[36mFile \u001B[39m\u001B[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/optim/adam.py:248\u001B[39m, in \u001B[36mAdam.step\u001B[39m\u001B[34m(self, closure)\u001B[39m\n\u001B[32m 236\u001B[39m beta1, beta2 = group[\u001B[33m\"\u001B[39m\u001B[33mbetas\u001B[39m\u001B[33m\"\u001B[39m]\n\u001B[32m 238\u001B[39m has_complex = \u001B[38;5;28mself\u001B[39m._init_group(\n\u001B[32m 239\u001B[39m group,\n\u001B[32m 240\u001B[39m params_with_grad,\n\u001B[32m (...)\u001B[39m\u001B[32m 245\u001B[39m state_steps,\n\u001B[32m 246\u001B[39m )\n\u001B[32m--> \u001B[39m\u001B[32m248\u001B[39m \u001B[43madam\u001B[49m\u001B[43m(\u001B[49m\n\u001B[32m 249\u001B[39m \u001B[43m \u001B[49m\u001B[43mparams_with_grad\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 250\u001B[39m \u001B[43m \u001B[49m\u001B[43mgrads\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 251\u001B[39m \u001B[43m \u001B[49m\u001B[43mexp_avgs\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 252\u001B[39m \u001B[43m \u001B[49m\u001B[43mexp_avg_sqs\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 253\u001B[39m \u001B[43m \u001B[49m\u001B[43mmax_exp_avg_sqs\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 254\u001B[39m \u001B[43m \u001B[49m\u001B[43mstate_steps\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 255\u001B[39m \u001B[43m \u001B[49m\u001B[43mamsgrad\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mamsgrad\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 256\u001B[39m \u001B[43m \u001B[49m\u001B[43mhas_complex\u001B[49m\u001B[43m=\u001B[49m\u001B[43mhas_complex\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 257\u001B[39m \u001B[43m \u001B[49m\u001B[43mbeta1\u001B[49m\u001B[43m=\u001B[49m\u001B[43mbeta1\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 258\u001B[39m \u001B[43m \u001B[49m\u001B[43mbeta2\u001B[49m\u001B[43m=\u001B[49m\u001B[43mbeta2\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 259\u001B[39m \u001B[43m \u001B[49m\u001B[43mlr\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mlr\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 260\u001B[39m \u001B[43m \u001B[49m\u001B[43mweight_decay\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mweight_decay\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 261\u001B[39m \u001B[43m \u001B[49m\u001B[43meps\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43meps\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 262\u001B[39m \u001B[43m \u001B[49m\u001B[43mmaximize\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mmaximize\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 263\u001B[39m \u001B[43m \u001B[49m\u001B[43mforeach\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mforeach\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 264\u001B[39m \u001B[43m \u001B[49m\u001B[43mcapturable\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mcapturable\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 265\u001B[39m \u001B[43m \u001B[49m\u001B[43mdifferentiable\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mdifferentiable\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 266\u001B[39m \u001B[43m \u001B[49m\u001B[43mfused\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mfused\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 267\u001B[39m \u001B[43m \u001B[49m\u001B[43mgrad_scale\u001B[49m\u001B[43m=\u001B[49m\u001B[38;5;28;43mgetattr\u001B[39;49m\u001B[43m(\u001B[49m\u001B[38;5;28;43mself\u001B[39;49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mgrad_scale\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[38;5;28;43;01mNone\u001B[39;49;00m\u001B[43m)\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 268\u001B[39m \u001B[43m \u001B[49m\u001B[43mfound_inf\u001B[49m\u001B[43m=\u001B[49m\u001B[38;5;28;43mgetattr\u001B[39;49m\u001B[43m(\u001B[49m\u001B[38;5;28;43mself\u001B[39;49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mfound_inf\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[38;5;28;43;01mNone\u001B[39;49;00m\u001B[43m)\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 269\u001B[39m \u001B[43m \u001B[49m\u001B[43mdecoupled_weight_decay\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgroup\u001B[49m\u001B[43m[\u001B[49m\u001B[33;43m\"\u001B[39;49m\u001B[33;43mdecoupled_weight_decay\u001B[39;49m\u001B[33;43m\"\u001B[39;49m\u001B[43m]\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 270\u001B[39m \u001B[43m \u001B[49m\u001B[43m)\u001B[49m\n\u001B[32m 272\u001B[39m \u001B[38;5;28;01mreturn\u001B[39;00m loss\n", "\u001B[36mFile \u001B[39m\u001B[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/optim/optimizer.py:151\u001B[39m, in \u001B[36m_disable_dynamo_if_unsupported..wrapper..maybe_fallback\u001B[39m\u001B[34m(*args, **kwargs)\u001B[39m\n\u001B[32m 149\u001B[39m \u001B[38;5;28;01mreturn\u001B[39;00m disabled_func(*args, **kwargs)\n\u001B[32m 150\u001B[39m \u001B[38;5;28;01melse\u001B[39;00m:\n\u001B[32m--> \u001B[39m\u001B[32m151\u001B[39m \u001B[38;5;28;01mreturn\u001B[39;00m \u001B[43mfunc\u001B[49m\u001B[43m(\u001B[49m\u001B[43m*\u001B[49m\u001B[43margs\u001B[49m\u001B[43m,\u001B[49m\u001B[43m \u001B[49m\u001B[43m*\u001B[49m\u001B[43m*\u001B[49m\u001B[43mkwargs\u001B[49m\u001B[43m)\u001B[49m\n", "\u001B[36mFile \u001B[39m\u001B[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/optim/adam.py:970\u001B[39m, in \u001B[36madam\u001B[39m\u001B[34m(params, grads, exp_avgs, exp_avg_sqs, max_exp_avg_sqs, state_steps, foreach, capturable, differentiable, fused, grad_scale, found_inf, has_complex, decoupled_weight_decay, amsgrad, beta1, beta2, lr, weight_decay, eps, maximize)\u001B[39m\n\u001B[32m 967\u001B[39m \u001B[38;5;28;01melse\u001B[39;00m:\n\u001B[32m 968\u001B[39m func = _single_tensor_adam\n\u001B[32m--> \u001B[39m\u001B[32m970\u001B[39m \u001B[43mfunc\u001B[49m\u001B[43m(\u001B[49m\n\u001B[32m 971\u001B[39m \u001B[43m \u001B[49m\u001B[43mparams\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 972\u001B[39m \u001B[43m \u001B[49m\u001B[43mgrads\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 973\u001B[39m \u001B[43m \u001B[49m\u001B[43mexp_avgs\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 974\u001B[39m \u001B[43m \u001B[49m\u001B[43mexp_avg_sqs\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 975\u001B[39m \u001B[43m \u001B[49m\u001B[43mmax_exp_avg_sqs\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 976\u001B[39m \u001B[43m \u001B[49m\u001B[43mstate_steps\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 977\u001B[39m \u001B[43m \u001B[49m\u001B[43mamsgrad\u001B[49m\u001B[43m=\u001B[49m\u001B[43mamsgrad\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 978\u001B[39m \u001B[43m \u001B[49m\u001B[43mhas_complex\u001B[49m\u001B[43m=\u001B[49m\u001B[43mhas_complex\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 979\u001B[39m \u001B[43m \u001B[49m\u001B[43mbeta1\u001B[49m\u001B[43m=\u001B[49m\u001B[43mbeta1\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 980\u001B[39m \u001B[43m \u001B[49m\u001B[43mbeta2\u001B[49m\u001B[43m=\u001B[49m\u001B[43mbeta2\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 981\u001B[39m \u001B[43m \u001B[49m\u001B[43mlr\u001B[49m\u001B[43m=\u001B[49m\u001B[43mlr\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 982\u001B[39m \u001B[43m \u001B[49m\u001B[43mweight_decay\u001B[49m\u001B[43m=\u001B[49m\u001B[43mweight_decay\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 983\u001B[39m \u001B[43m \u001B[49m\u001B[43meps\u001B[49m\u001B[43m=\u001B[49m\u001B[43meps\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 984\u001B[39m \u001B[43m \u001B[49m\u001B[43mmaximize\u001B[49m\u001B[43m=\u001B[49m\u001B[43mmaximize\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 985\u001B[39m \u001B[43m \u001B[49m\u001B[43mcapturable\u001B[49m\u001B[43m=\u001B[49m\u001B[43mcapturable\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 986\u001B[39m \u001B[43m \u001B[49m\u001B[43mdifferentiable\u001B[49m\u001B[43m=\u001B[49m\u001B[43mdifferentiable\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 987\u001B[39m \u001B[43m \u001B[49m\u001B[43mgrad_scale\u001B[49m\u001B[43m=\u001B[49m\u001B[43mgrad_scale\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 988\u001B[39m \u001B[43m \u001B[49m\u001B[43mfound_inf\u001B[49m\u001B[43m=\u001B[49m\u001B[43mfound_inf\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 989\u001B[39m \u001B[43m \u001B[49m\u001B[43mdecoupled_weight_decay\u001B[49m\u001B[43m=\u001B[49m\u001B[43mdecoupled_weight_decay\u001B[49m\u001B[43m,\u001B[49m\n\u001B[32m 990\u001B[39m \u001B[43m\u001B[49m\u001B[43m)\u001B[49m\n", "\u001B[36mFile \u001B[39m\u001B[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/optim/adam.py:545\u001B[39m, in \u001B[36m_single_tensor_adam\u001B[39m\u001B[34m(params, grads, exp_avgs, exp_avg_sqs, max_exp_avg_sqs, state_steps, grad_scale, found_inf, amsgrad, has_complex, beta1, beta2, lr, weight_decay, eps, maximize, capturable, differentiable, decoupled_weight_decay)\u001B[39m\n\u001B[32m 543\u001B[39m denom = (max_exp_avg_sqs[i].sqrt() / bias_correction2_sqrt).add_(eps)\n\u001B[32m 544\u001B[39m \u001B[38;5;28;01melse\u001B[39;00m:\n\u001B[32m--> \u001B[39m\u001B[32m545\u001B[39m denom = (\u001B[43mexp_avg_sq\u001B[49m\u001B[43m.\u001B[49m\u001B[43msqrt\u001B[49m\u001B[43m(\u001B[49m\u001B[43m)\u001B[49m / bias_correction2_sqrt).add_(eps)\n\u001B[32m 547\u001B[39m param.addcdiv_(exp_avg, denom, value=-step_size) \u001B[38;5;66;03m# type: ignore[arg-type]\u001B[39;00m\n\u001B[32m 549\u001B[39m \u001B[38;5;66;03m# Lastly, switch back to complex view\u001B[39;00m\n", "\u001B[31mKeyboardInterrupt\u001B[39m: " ] } ], "execution_count": 58 }, { "cell_type": "code", "source": [ "class DatasetMnistTest(Dataset):\n", " def __init__(self,df):\n", " self.df=df\n", " def __len__(self):\n", " return len(self.df)\n", " def __getitem__(self, idx):\n", " item= {\"Data\": torch.tensor(self.df.iloc[idx].to_numpy().reshape(28, 28),dtype=torch.float)}\n", " return item\n", "test_dataset = DatasetMnistTest(test_df)\n", "test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)\n", "all_preds = []\n", "with torch.no_grad():\n", " for batch in test_loader:\n", " input_ids = batch['Data'].to(device)\n", " outputs = model(input_ids)\n", " preds = torch.argmax(outputs, dim=1)\n", " all_preds.extend(preds.cpu().numpy())\n", "idd = range(1,len(all_preds)+1)\n", "submission = pd.DataFrame({\n", " 'ImageId':idd,\n", " 'Label': all_preds\n", "})\n", "print(submission)\n", "submission.to_csv('submission.csv', index=False)\n", "print(\"Submission saved!\")" ], "metadata": { "trusted": true, "execution": { "iopub.status.busy": "2026-07-04T16:19:40.600448Z", "iopub.execute_input": "2026-07-04T16:19:40.601006Z", "iopub.status.idle": "2026-07-04T16:19:41.910494Z", "shell.execute_reply.started": "2026-07-04T16:19:40.600982Z", "shell.execute_reply": "2026-07-04T16:19:41.909862Z" } }, "outputs": [], "execution_count": null }, { "cell_type": "code", "source": "", "metadata": { "trusted": true }, "outputs": [], "execution_count": null } ] }