{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport torch\nimport torchvision\nimport torchvision.transforms as transforms\nimport torch.optim as optim\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision.io import read_image\nfrom torchvision.utils import make_grid\nfrom torch.utils.data import random_split, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom PIL import Image\nimport warnings\nwarnings.simplefilter(\"ignore\", category=DeprecationWarning)\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-06T19:50:01.498521Z","iopub.execute_input":"2022-12-06T19:50:01.498889Z","iopub.status.idle":"2022-12-06T19:50:01.509389Z","shell.execute_reply.started":"2022-12-06T19:50:01.498857Z","shell.execute_reply":"2022-12-06T19:50:01.508268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:01.512941Z","iopub.execute_input":"2022-12-06T19:50:01.513962Z","iopub.status.idle":"2022-12-06T19:50:01.519653Z","shell.execute_reply.started":"2022-12-06T19:50:01.513925Z","shell.execute_reply":"2022-12-06T19:50:01.518605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_path = '../input/histopathologic-cancer-detection/test'\ntrain_path = '../input/histopathologic-cancer-detection/train'\nsample_submission = pd.read_csv('../input/histopathologic-cancer-detection/sample_submission.csv')\ntrain_labels = pd.read_csv('../input/histopathologic-cancer-detection/train_labels.csv')","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:01.523934Z","iopub.execute_input":"2022-12-06T19:50:01.524231Z","iopub.status.idle":"2022-12-06T19:50:01.742019Z","shell.execute_reply.started":"2022-12-06T19:50:01.524184Z","shell.execute_reply":"2022-12-06T19:50:01.740878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:01.744193Z","iopub.execute_input":"2022-12-06T19:50:01.744687Z","iopub.status.idle":"2022-12-06T19:50:01.755125Z","shell.execute_reply.started":"2022-12-06T19:50:01.744650Z","shell.execute_reply":"2022-12-06T19:50:01.753946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels['label'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:01.757053Z","iopub.execute_input":"2022-12-06T19:50:01.757483Z","iopub.status.idle":"2022-12-06T19:50:01.770238Z","shell.execute_reply.started":"2022-12-06T19:50:01.757449Z","shell.execute_reply":"2022-12-06T19:50:01.768872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.pie(train_labels['label'].value_counts(), labels=[0,1])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:01.772920Z","iopub.execute_input":"2022-12-06T19:50:01.773296Z","iopub.status.idle":"2022-12-06T19:50:01.853733Z","shell.execute_reply.started":"2022-12-06T19:50:01.773261Z","shell.execute_reply":"2022-12-06T19:50:01.847923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.countplot(train_labels['label'])","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:01.855146Z","iopub.execute_input":"2022-12-06T19:50:01.855536Z","iopub.status.idle":"2022-12-06T19:50:02.133652Z","shell.execute_reply.started":"2022-12-06T19:50:01.855498Z","shell.execute_reply":"2022-12-06T19:50:02.127715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cancerous_img = train_labels.loc[train_labels['label']==1]\nhealthy_img = train_labels.loc[train_labels['label']==0]","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:02.134904Z","iopub.execute_input":"2022-12-06T19:50:02.135516Z","iopub.status.idle":"2022-12-06T19:50:02.159070Z","shell.execute_reply.started":"2022-12-06T19:50:02.135476Z","shell.execute_reply":"2022-12-06T19:50:02.158133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cancerous_img.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:02.160225Z","iopub.execute_input":"2022-12-06T19:50:02.160767Z","iopub.status.idle":"2022-12-06T19:50:02.174271Z","shell.execute_reply.started":"2022-12-06T19:50:02.160732Z","shell.execute_reply":"2022-12-06T19:50:02.172254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"healthy_img.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:02.178003Z","iopub.execute_input":"2022-12-06T19:50:02.180192Z","iopub.status.idle":"2022-12-06T19:50:02.193403Z","shell.execute_reply.started":"2022-12-06T19:50:02.180159Z","shell.execute_reply":"2022-12-06T19:50:02.192576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nfig, ax = plt.subplots(5, 10)\nfor i, j in enumerate(cancerous_img['id'][:50]):\n    img_path = os.path.join(train_path,j+'.tif')\n    img = Image.open(img_path)\n    plt.subplot(5,10,i+1)\n    plt.imshow(np.array(img))\n    plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:02.197190Z","iopub.execute_input":"2022-12-06T19:50:02.197772Z","iopub.status.idle":"2022-12-06T19:50:04.557179Z","shell.execute_reply.started":"2022-12-06T19:50:02.197739Z","shell.execute_reply":"2022-12-06T19:50:04.556157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nfig, ax = plt.subplots(5, 10)\nfor i, j in enumerate(healthy_img['id'][:50]):\n    img_path = os.path.join(train_path,j+'.tif')\n    img = Image.open(img_path)\n    plt.subplot(5,10,i+1)\n    plt.imshow(np.array(img))\n    plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:04.562165Z","iopub.execute_input":"2022-12-06T19:50:04.562954Z","iopub.status.idle":"2022-12-06T19:50:06.137991Z","shell.execute_reply.started":"2022-12-06T19:50:04.562912Z","shell.execute_reply":"2022-12-06T19:50:06.137151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = train_labels['label']","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:06.139621Z","iopub.execute_input":"2022-12-06T19:50:06.140271Z","iopub.status.idle":"2022-12-06T19:50:06.144845Z","shell.execute_reply.started":"2022-12-06T19:50:06.140223Z","shell.execute_reply":"2022-12-06T19:50:06.143902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataLoader():\n    def __init__(self, labels, dir_path, transforms=None, show=False):\n        self.labels = labels\n        self.dir_path = dir_path\n        self.transform = transforms\n        self.show=show\n        \n    def __len__(self):\n        return len(self.labels)\n    \n    def __gettransforms__(self):\n        return self.transform\n    \n    def __getshow__(self):\n        return self.show\n    \n    def __getitem__(self, idx):\n        img_path = os.path.join(self.dir_path, self.labels.iloc[idx,0] +'.tif')\n        image = Image.open(img_path)\n        label = self.labels.iloc[idx,1]\n        \n        if self.transform:\n            image = self.transform(image)\n            if self.show == True:\n                plt.imshow(image.permute(1,2,0))\n        else:\n            if self.show == True:\n                plt.imshow(image)\n        return image, label\n        ","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:06.146092Z","iopub.execute_input":"2022-12-06T19:50:06.147609Z","iopub.status.idle":"2022-12-06T19:50:06.157722Z","shell.execute_reply.started":"2022-12-06T19:50:06.147575Z","shell.execute_reply":"2022-12-06T19:50:06.156671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(25),\n    transforms.Grayscale(),\n    transforms.GaussianBlur((7,13)),\n    transforms.ToTensor()])","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:06.158772Z","iopub.execute_input":"2022-12-06T19:50:06.159036Z","iopub.status.idle":"2022-12-06T19:50:06.171711Z","shell.execute_reply.started":"2022-12-06T19:50:06.159013Z","shell.execute_reply":"2022-12-06T19:50:06.170704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = CustomDataLoader(train_labels, train_path, show=True)\ndata_transform = CustomDataLoader(train_labels, train_path, transforms=transform, show=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:06.173894Z","iopub.execute_input":"2022-12-06T19:50:06.174894Z","iopub.status.idle":"2022-12-06T19:50:06.181619Z","shell.execute_reply.started":"2022-12-06T19:50:06.174857Z","shell.execute_reply":"2022-12-06T19:50:06.180614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(data.__gettransforms__(), data_transform.__gettransforms__())","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:06.184344Z","iopub.execute_input":"2022-12-06T19:50:06.184956Z","iopub.status.idle":"2022-12-06T19:50:06.191647Z","shell.execute_reply.started":"2022-12-06T19:50:06.184923Z","shell.execute_reply":"2022-12-06T19:50:06.190579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random_idx = [np.random.randint(1,data.__len__()) for _ in range(10)]\nprint(random_idx)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:06.193283Z","iopub.execute_input":"2022-12-06T19:50:06.193656Z","iopub.status.idle":"2022-12-06T19:50:06.201797Z","shell.execute_reply.started":"2022-12-06T19:50:06.193623Z","shell.execute_reply":"2022-12-06T19:50:06.200672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nfor i,j in enumerate(random_idx):\n    plt.subplot(2,5,i+1)\n    data.__getitem__(j)\n    plt.title('Random Non-Augmented Images')\n    plt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:06.203616Z","iopub.execute_input":"2022-12-06T19:50:06.204037Z","iopub.status.idle":"2022-12-06T19:50:06.967418Z","shell.execute_reply.started":"2022-12-06T19:50:06.204004Z","shell.execute_reply":"2022-12-06T19:50:06.966528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nfor i,j in enumerate(random_idx):\n    plt.subplot(2,5,i+1)\n    data_transform.__getitem__(j)\n    plt.title('Random Augmented Images')\n    plt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:06.968866Z","iopub.execute_input":"2022-12-06T19:50:06.969470Z","iopub.status.idle":"2022-12-06T19:50:07.676933Z","shell.execute_reply.started":"2022-12-06T19:50:06.969435Z","shell.execute_reply":"2022-12-06T19:50:07.676094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## PART 3: BUILDING THE ACTUAL DATASET","metadata":{}},{"cell_type":"code","source":"data =  CustomDataLoader(train_labels, train_path)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:07.678407Z","iopub.execute_input":"2022-12-06T19:50:07.679366Z","iopub.status.idle":"2022-12-06T19:50:07.683930Z","shell.execute_reply.started":"2022-12-06T19:50:07.679329Z","shell.execute_reply":"2022-12-06T19:50:07.683093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_transforms = transforms.Compose([\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(25),\n    transforms.ToTensor(),\n    transforms.Normalize((0.5,0.5,0.5), (0.5,0.5,0.5))\n])","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:07.685180Z","iopub.execute_input":"2022-12-06T19:50:07.686186Z","iopub.status.idle":"2022-12-06T19:50:07.693545Z","shell.execute_reply.started":"2022-12-06T19:50:07.686152Z","shell.execute_reply":"2022-12-06T19:50:07.692705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = train_labels[:10000]\ntrain_labels.label.value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:07.695233Z","iopub.execute_input":"2022-12-06T19:50:07.696389Z","iopub.status.idle":"2022-12-06T19:50:07.710412Z","shell.execute_reply.started":"2022-12-06T19:50:07.696355Z","shell.execute_reply":"2022-12-06T19:50:07.709368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train, val = train_test_split(train_labels, stratify=train_labels.label, test_size=0.2)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:07.711892Z","iopub.execute_input":"2022-12-06T19:50:07.713103Z","iopub.status.idle":"2022-12-06T19:50:07.724463Z","shell.execute_reply.started":"2022-12-06T19:50:07.713068Z","shell.execute_reply":"2022-12-06T19:50:07.723543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = CustomDataLoader(train, train_path, training_transforms, show=True)\nval_data = CustomDataLoader(val, train_path, transforms.Compose([transforms.ToTensor(),\n                                                                transforms.Normalize((0.5,0.5,0.5), (0.5,0.5,0.5))]),\n                           show=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:07.725648Z","iopub.execute_input":"2022-12-06T19:50:07.726506Z","iopub.status.idle":"2022-12-06T19:50:07.732412Z","shell.execute_reply.started":"2022-12-06T19:50:07.726473Z","shell.execute_reply":"2022-12-06T19:50:07.731357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.__len__(), val_data.__len__()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:07.733900Z","iopub.execute_input":"2022-12-06T19:50:07.734864Z","iopub.status.idle":"2022-12-06T19:50:07.743744Z","shell.execute_reply.started":"2022-12-06T19:50:07.734831Z","shell.execute_reply":"2022-12-06T19:50:07.742718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nfor i in range(10):\n    plt.subplot(2,5,i+1)\n    train_data.__getitem__(i)\n    plt.title('Random Augmented Training images')\n    plt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:07.745294Z","iopub.execute_input":"2022-12-06T19:50:07.745904Z","iopub.status.idle":"2022-12-06T19:50:08.410154Z","shell.execute_reply.started":"2022-12-06T19:50:07.745838Z","shell.execute_reply":"2022-12-06T19:50:08.409285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nfor i in range(10):\n    plt.subplot(2,5,i+1)\n    val_data.__getitem__(i)\n    plt.title('Random Non-Augmented Validation images')\n    plt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:08.411490Z","iopub.execute_input":"2022-12-06T19:50:08.412443Z","iopub.status.idle":"2022-12-06T19:50:09.603994Z","shell.execute_reply.started":"2022-12-06T19:50:08.412407Z","shell.execute_reply":"2022-12-06T19:50:09.599708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.show=False\nval_data.show=False","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:09.605607Z","iopub.execute_input":"2022-12-06T19:50:09.606227Z","iopub.status.idle":"2022-12-06T19:50:09.611472Z","shell.execute_reply.started":"2022-12-06T19:50:09.606176Z","shell.execute_reply":"2022-12-06T19:50:09.610642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_data, batch_size=256)\nval_loader = DataLoader(val_data, batch_size=256)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:09.618991Z","iopub.execute_input":"2022-12-06T19:50:09.619380Z","iopub.status.idle":"2022-12-06T19:50:09.624834Z","shell.execute_reply.started":"2022-12-06T19:50:09.619346Z","shell.execute_reply":"2022-12-06T19:50:09.623960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## PART 4: BUILDING THE MODEL\n\n### Architecture: I've tried two architectures: the first one (the actual model we will be using) will use two 2D convolution layers, which doubles the number of channels at each step. Then we have a maxpooling layer to reduce the size of the image. We again have two 2D convolution layers, which ends with 48 output channels. Finally, the last layer is a fully connected layer that we will use to make our predictions. The final size of the image is 48 * 24 * 24.","metadata":{}},{"cell_type":"code","source":"class CNN_1(nn.Module):\n    def __init__(self, num_classes = 2):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels=3, out_channels=6, kernel_size=3, stride=(1,1), padding=(1,1))\n        self.conv2 = nn.Conv2d(in_channels=6, out_channels=12, kernel_size=3, stride=(1,1), padding=(1,1))\n        self.pool = nn.MaxPool2d(kernel_size=(2,2), stride=(2,2)) # Half the dimension size\n        self.conv3 = nn.Conv2d(in_channels=12, out_channels=24, kernel_size=3, stride=(1,1), padding=(1,1))\n        self.conv4 = nn.Conv2d(in_channels=24, out_channels=48, kernel_size=3, stride=(1,1), padding=(1,1))\n        self.fc = nn.Linear(48*24*24, num_classes)\n        self.sigmoid = nn.Sigmoid()\n        \n    def forward(self,x):\n        x = F.relu(self.conv1(x))\n        x = F.relu(self.conv2(x))\n        x = self.pool(x)\n        x = F.relu(self.conv3(x))\n        x = F.relu(self.conv4(x))\n        x = self.pool(x)\n        x = x.view(-1, 48*24*24)\n        x = self.fc(x)\n        x = self.sigmoid(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:53:21.752024Z","iopub.execute_input":"2022-12-06T19:53:21.753167Z","iopub.status.idle":"2022-12-06T19:53:21.764177Z","shell.execute_reply.started":"2022-12-06T19:53:21.753119Z","shell.execute_reply":"2022-12-06T19:53:21.763214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#class CNN_2(nn.Module):\n#    def __init__(self, num_classes = 2):\n#        super().__init__()\n#        self.conv1 = nn.Conv2d(in_channels=3, out_channels=6, kernel_size=3, stride=(1,1), padding=(1,1))\n#        self.conv2 = nn.Conv2d(in_channels=6, out_channels=12, kernel_size=3, stride=(1,1), padding=(1,1))\n#        self.pool = nn.MaxPool2d(kernel_size=(2,2), stride=(2,2)) # Half the dimension size\n#        self.fc = nn.Linear(12*48*48, num_classes)\n#        \n#    def forward(self,x):\n#        x = F.relu(self.conv1(x))\n#        x = F.relu(self.conv2(x))        \n#        x = self.pool(x)\n#        x = x.view(-1, 12*48*48)\n#        x = self.fc(x)\n#        return x","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:53:22.303496Z","iopub.execute_input":"2022-12-06T19:53:22.303841Z","iopub.status.idle":"2022-12-06T19:53:22.308776Z","shell.execute_reply.started":"2022-12-06T19:53:22.303812Z","shell.execute_reply":"2022-12-06T19:53:22.307774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### For the Loss function, we use CrossEntropy. We could have used the negative log likelihood loss as well, however for our case the CrossEntropy had better results. Furthermore, we use Adam as the optimizer since it's the fastest and most reliable. At each step of the training, we will be manually computing the training loss and accuracy so that we can plot it (we could have used tensorboard but Kaggle doesn't support it)","metadata":{}},{"cell_type":"code","source":"model_1 = CNN_1().to(device)\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adamax(model_1.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:53:23.187083Z","iopub.execute_input":"2022-12-06T19:53:23.189521Z","iopub.status.idle":"2022-12-06T19:53:23.197953Z","shell.execute_reply.started":"2022-12-06T19:53:23.189482Z","shell.execute_reply":"2022-12-06T19:53:23.196878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loss_model_1 = []\ntrain_acc_model_1 = []\ncorrect_model_1 = 0\ntotal_model_1 = 0\nsteps = len(train_loader)\n\nfor epoch in range(5):\n    current_loss_model_1 = 0\n    for i, (images, labels) in enumerate(train_loader):\n        images = images.to(device)\n        labels = labels.to(device)\n        \n        # Forward pass\n        outputs = model_1(images)\n        loss = criterion(outputs, labels)\n        \n        # Backwards pass\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        current_loss_model_1 += loss.item()\n        _, pred = outputs.max(1)\n        total_model_1 += labels.shape[0]\n        correct_model_1 += pred.eq(labels).sum().item()\n        \n        if (i+1) % 10 == 0:\n            print(f'Epoch: {epoch+1}/{5}, Step: {i+1}/{steps}, Loss: {loss.item()}')\n    \n    acc = 100.0 * correct_model_1/total_model_1\n    train_loss_model_1.append(current_loss_model_1/steps)\n    train_acc_model_1.append(acc)\n    print(f'The accuracy for the epoch {epoch+1} is : {acc}')\n\nprint('END')","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:53:23.830873Z","iopub.execute_input":"2022-12-06T19:53:23.831241Z","iopub.status.idle":"2022-12-06T19:53:33.851012Z","shell.execute_reply.started":"2022-12-06T19:53:23.831187Z","shell.execute_reply":"2022-12-06T19:53:33.849400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nplt.plot([1,2,3,4,5], train_acc_model_1)\nplt.xlabel('Epochs')\nplt.ylabel('Accuracy (in percentages)')\nplt.title('Accuracy against epochs')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:10.396074Z","iopub.status.idle":"2022-12-06T19:50:10.396923Z","shell.execute_reply.started":"2022-12-06T19:50:10.396658Z","shell.execute_reply":"2022-12-06T19:50:10.396683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nplt.plot([1,2,3,4,5], train_loss_model_1)\nplt.xlabel('Epochs')\nplt.ylabel('Loss (in percentages)')\nplt.title('Loss against epochs')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:10.398495Z","iopub.status.idle":"2022-12-06T19:50:10.399315Z","shell.execute_reply.started":"2022-12-06T19:50:10.399026Z","shell.execute_reply":"2022-12-06T19:50:10.399051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check accuracy on the validation dataset","metadata":{}},{"cell_type":"code","source":"with torch.no_grad():\n    total, correct = 0,0\n    for images, labels in val_loader:\n        images, labels = images.to(device), labels.to(device)\n        out = model_1(images)\n        _, preds = out.max(1)\n        total += labels.shape[0]\n        correct += (preds==labels).sum().item()\n    print(f'The accuracy on the validation set is: {100*correct/total}%')","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:10.400735Z","iopub.status.idle":"2022-12-06T19:50:10.401602Z","shell.execute_reply.started":"2022-12-06T19:50:10.401340Z","shell.execute_reply":"2022-12-06T19:50:10.401367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## FINAL PART: MAKING THE PREDICTIONS","metadata":{}},{"cell_type":"code","source":"test_data = CustomDataLoader(sample_submission, test_path, \n                             transforms=transforms.Compose([transforms.ToTensor(),\n                                                                transforms.Normalize((0.5,0.5,0.5), (0.5,0.5,0.5))]))","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:10.403060Z","iopub.status.idle":"2022-12-06T19:50:10.403862Z","shell.execute_reply.started":"2022-12-06T19:50:10.403594Z","shell.execute_reply":"2022-12-06T19:50:10.403620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader=DataLoader(test_data, \n                       batch_size=256)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:10.405291Z","iopub.status.idle":"2022-12-06T19:50:10.406107Z","shell.execute_reply.started":"2022-12-06T19:50:10.405848Z","shell.execute_reply":"2022-12-06T19:50:10.405874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_1.eval()\n\npredictions = []\n\nfor i, (images, labels) in enumerate(test_loader):\n    images, labels = images.to(device), labels.to(device)\n    out = model_1(images)\n    \n    preds = out[:,1].detach().cpu().numpy()\n    for i in preds:\n        predictions.append(i)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:10.407557Z","iopub.status.idle":"2022-12-06T19:50:10.408453Z","shell.execute_reply.started":"2022-12-06T19:50:10.408172Z","shell.execute_reply":"2022-12-06T19:50:10.408211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(predictions[:5])","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:10.409936Z","iopub.status.idle":"2022-12-06T19:50:10.410745Z","shell.execute_reply.started":"2022-12-06T19:50:10.410486Z","shell.execute_reply":"2022-12-06T19:50:10.410512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission['label'] = predictions\nsample_submission.head()\nsample_submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-12-06T19:50:10.412194Z","iopub.status.idle":"2022-12-06T19:50:10.412971Z","shell.execute_reply.started":"2022-12-06T19:50:10.412710Z","shell.execute_reply":"2022-12-06T19:50:10.412735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}