{"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":"markdown","source":"# Reason Why This Notebook Exists","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-06-23T19:30:21.918141Z","iopub.execute_input":"2021-06-23T19:30:21.918535Z","iopub.status.idle":"2021-06-23T19:30:21.927489Z","shell.execute_reply.started":"2021-06-23T19:30:21.918450Z","shell.execute_reply":"2021-06-23T19:30:21.926653Z"}}},{"cell_type":"markdown","source":"### There are various notebooks on Kaggle describing how to get started with the image classificatin using pyutorch or Tensorflow to detect Covid-19 by analysing X-rays. However most of them are highly complicated or least they are complicated to me. So I decided to help out everyone by creating this ver simple starte notebook. It shows how you can train an EfficientNet model using the Kaggle data set and pytorch library.","metadata":{}},{"cell_type":"markdown","source":"# Introduction TO Covid-19 problem","metadata":{}},{"cell_type":"markdown","source":"# Importing Libraries","metadata":{}},{"cell_type":"code","source":"# Libraries and dependencies needed for dicom visualation and model training\n# !pip install torchsummary\n# !pip install efficientnet_pytorch\n# !pip install pylibjpeg pylibjpeg-libjpeg\n# !conda install -c conda-forge gdcm -y\n# !pip install python-gdcm\n# !pip install pylibjpeg\n# import pydicom\n# !python -m pip uninstall numpy --yes\n# !pip install numpy==1.19.2\n","metadata":{"execution":{"iopub.status.busy":"2021-06-19T23:29:41.083854Z","iopub.execute_input":"2021-06-19T23:29:41.08421Z","iopub.status.idle":"2021-06-19T23:30:13.761888Z","shell.execute_reply.started":"2021-06-19T23:29:41.084179Z","shell.execute_reply":"2021-06-19T23:30:13.76033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\n!pip install efficientnet_pytorch\n!pip install torchsummary\nfrom torchsummary import summary\nimport PIL\nimport sys\nimport torch\nfrom time import time\nimport torchvision\nfrom PIL import Image\nimport torch.nn as nn\nfrom torch.utils import data\nfrom torch.autograd import Variable\nimport torchvision.transforms as transforms\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\n# I have used Efficientnet3 pre-trained model for the classification task\nfrom efficientnet_pytorch import EfficientNet\n# !python -m pip uninstall numpy --yes\n# !python -m pip uninstall pydicom --yes\n# !pip install pydicom\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\n\nfrom torch.autograd import Variable\nimport torchvision.transforms as transforms\nimport matplotlib.pyplot as plt\nimport torch.nn.functional as F\nimport torch.optim as om\nimport torchvision as tv\nimport torch.utils.data as dat\n\n# I have used Efficientnet3 pre-trained model for the classification task\n\nif torch.cuda.is_available():     # Make sure GPU is available\n    dev = torch.device(\"cuda:0\")\n    kwar = {'num_workers': 8, 'pin_memory': True}\n    cpu = torch.device(\"cpu\")\nelse:\n    print(\"Warning: CUDA not found, CPU only.\")\n    dev = torch.device(\"cpu\")\n    kwar = {}\n    cpu = torch.device(\"cpu\")","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:30:21.929009Z","iopub.execute_input":"2021-06-23T19:30:21.929502Z","iopub.status.idle":"2021-06-23T19:30:40.144975Z","shell.execute_reply.started":"2021-06-23T19:30:21.929461Z","shell.execute_reply":"2021-06-23T19:30:40.142536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Learning About Datasets","metadata":{}},{"cell_type":"markdown","source":"## Loading the training set image and study ids","metadata":{}},{"cell_type":"code","source":"training_set_study = pd.read_csv('../input/siim-covid19-detection/train_study_level.csv')\ntraining_set_image_level = pd.read_csv('../input/siim-covid19-detection/train_image_level.csv')\ntraining_set_study.head(100)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:30:54.432981Z","iopub.execute_input":"2021-06-23T19:30:54.433333Z","iopub.status.idle":"2021-06-23T19:30:54.514458Z","shell.execute_reply.started":"2021-06-23T19:30:54.433285Z","shell.execute_reply":"2021-06-23T19:30:54.513661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# looking at the image_level_training_set.\ntraining_set_image_level.head(100)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:30:57.025219Z","iopub.execute_input":"2021-06-23T19:30:57.025569Z","iopub.status.idle":"2021-06-23T19:30:57.037933Z","shell.execute_reply.started":"2021-06-23T19:30:57.025539Z","shell.execute_reply":"2021-06-23T19:30:57.037156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Changing the column values to allow for merging of two tables. The merged tables can then be used to extract image data.\ntraining_set_image_level.id = training_set_image_level.id.str.replace('_image','')\ntraining_set_image_level.head()\ntraining_set_study.id = training_set_study.id.str.replace('_study','')\ntraining_set_study.head()\n# training_set_image_level.loc[training_set_image_level.StudyInstanceUID == '005057b3f880']","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:31:00.550910Z","iopub.execute_input":"2021-06-23T19:31:00.551226Z","iopub.status.idle":"2021-06-23T19:31:00.575339Z","shell.execute_reply.started":"2021-06-23T19:31:00.551196Z","shell.execute_reply":"2021-06-23T19:31:00.574400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Merging Two Training Set Tables\nThe StudyInstanceUID in train_image_level is same as the 'id' in train_study_level. Henceforth these two tables can be merged based on these two columns.","metadata":{}},{"cell_type":"code","source":"combined_training_set = training_set_image_level.merge(training_set_study, left_on = 'StudyInstanceUID', right_on = 'id', how = 'inner', suffixes=('_image', 'study'))\ncombined_training_set.drop(columns = 'idstudy', inplace = True)\ncombined_training_set.head()\n","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:31:02.990040Z","iopub.execute_input":"2021-06-23T19:31:02.990375Z","iopub.status.idle":"2021-06-23T19:31:03.018098Z","shell.execute_reply.started":"2021-06-23T19:31:02.990341Z","shell.execute_reply":"2021-06-23T19:31:03.017358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Getting path of each image and adding it to the combined tables\n<!--     Assumptions -->\n1. Some image folders have multiple images however we have neglected that and instead only selected the first image in that folder.","metadata":{}},{"cell_type":"code","source":"# Creating function to get absolute path.\n\ndef get_absolute_file_paths(x):\n    path = '../input/siim-covid19-detection/train/'\n    directory = os.path.join(path,x )\n    all_abs_file_paths = []\n    for dirpath,_,filenames in os.walk(directory):\n        for f in filenames:\n            all_abs_file_paths.append(os.path.abspath(os.path.join(dirpath, f)))\n    return all_abs_file_paths\n# get_absolute_file_paths(combined_training_set.StudyInstanceUID[0])","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:31:09.777977Z","iopub.execute_input":"2021-06-23T19:31:09.778295Z","iopub.status.idle":"2021-06-23T19:31:09.785468Z","shell.execute_reply.started":"2021-06-23T19:31:09.778264Z","shell.execute_reply":"2021-06-23T19:31:09.784517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"combined_training_set['abs_image_path'] = combined_training_set.StudyInstanceUID.apply(get_absolute_file_paths) \ncombined_training_set.head()","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:31:18.619806Z","iopub.execute_input":"2021-06-23T19:31:18.620135Z","iopub.status.idle":"2021-06-23T19:31:41.381760Z","shell.execute_reply.started":"2021-06-23T19:31:18.620105Z","shell.execute_reply":"2021-06-23T19:31:41.380899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reading X-rays for image visualisation","metadata":{}},{"cell_type":"markdown","source":"## Function for reading dicom images","metadata":{}},{"cell_type":"code","source":"# Function for reading .dicom files\n# ----------> Copid from 'siim-cov19 efnb7+yolov5 [infer]' :) ~~~~~~~~ Thanks\ndef read_xray(path, voi_lut = True, fix_monochrome = True):\n    # Original from: https://www.kaggle.com/raddar/convert-dicom-to-np-array-the-correct-way\n    dicom = pydicom.read_file(path)\n    \n    # VOI LUT (if available by DICOM device) is used to transform raw DICOM data to \n    # \"human-friendly\" view\n    if voi_lut:\n        data = apply_voi_lut(dicom.pixel_array, dicom)\n    else:\n        data = dicom.pixel_array\n               \n    # depending on this value, X-ray may look inverted - fix that:\n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = np.amax(data) - data\n        \n    data = data - np.min(data)\n    data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n        \n    return data","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:31:41.383217Z","iopub.execute_input":"2021-06-23T19:31:41.383597Z","iopub.status.idle":"2021-06-23T19:31:41.390215Z","shell.execute_reply.started":"2021-06-23T19:31:41.383560Z","shell.execute_reply":"2021-06-23T19:31:41.388949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualisation\n### Reading 10 random images and showing there labels","metadata":{}},{"cell_type":"code","source":"!pip install python-gdcm\n!pip install pylibjpeg pylibjpeg-libjpeg\nfig, axs = plt.subplots(nrows = 2, ncols = 5, sharex = True, sharey = True, figsize = (20, 10), tight_layout = True)\nfor ax, Index_num in zip(axs.reshape(-1), np.random.choice(range(len(combined_training_set)), 20, replace=False)):\n    data = read_xray(combined_training_set.abs_image_path[Index_num][0])\n    data = cv2.resize(data, (500,500))\n    ax.imshow(data, cmap = 'gray')\n    ax.set_title('{}'.format(combined_training_set[['Negative for Pneumonia', 'Typical Appearance',\n           'Indeterminate Appearance', 'Atypical Appearance']].idxmax(axis = 1)[Index_num]))","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:31:43.084942Z","iopub.execute_input":"2021-06-23T19:31:43.085283Z","iopub.status.idle":"2021-06-23T19:32:06.275662Z","shell.execute_reply.started":"2021-06-23T19:31:43.085252Z","shell.execute_reply":"2021-06-23T19:32:06.274785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Getting numpy array from dicom images\n- This numpy array can then be used as an input to any tensorflow or pytorch deep learning model","metadata":{}},{"cell_type":"code","source":"# currently ignoring the dicom image errors --------------> need help on that, please leave a comment if you can find a solution to it. Thanks in advance.\ndef getNumpyArraysSet(combined_training_set, ImageSize = (128,128)):\n    X_temp = []\n    Y_temp = []\n    dicom_error_images_path = []\n    i = 0\n    for n in range(len(combined_training_set)):\n        print('Percentage task complete {}'.format(n/len(combined_training_set) * 100))\n        try:\n            data = read_xray(combined_training_set.abs_image_path[n][0])\n            try:\n                Y_temp.append(combined_training_set[['Negative for Pneumonia', 'Typical Appearance',\n                   'Indeterminate Appearance', 'Atypical Appearance']].iloc[n].values)\n            except:\n                pass\n        except RuntimeError:\n            i = i+1\n            dicom_error_images_path.append(combined_training_set.abs_image_path[n][0])\n            print('The number of images not read because of pydicom library errors are {}'.format(i))\n    #     Get data and directory of images giving dicom error.\n            continue\n        data = cv2.resize(data, ImageSize)\n        X_temp.append(data)\n    # stacking up list in ndarray\n    print('stacking arrays in list started')\n    X = np.dstack(X_temp)\n    try:\n        y = np.dstack(Y_temp)\n    except:\n        print('y has no elemnts')\n    print('Stacking finished')\n    np.save('ImageArray'+str(ImageSize[0]), X)\n    try:\n        np.save('TargetArray'+str(ImageSize[0]), y)\n    except:\n        print('y has no elemnts')\n    # Saving path to images that were not analysed.....\n    df = pd.DataFrame(dicom_error_images_path)\n    df.to_csv('PathDicomErrorImages.csv')\n    print('All Done --------> Enjoy !!!!')\n    return X, y\n    \n##################### \n# Getting the numpy arrays for training data set. \nX, Y = getNumpyArraysSet(combined_training_set, ImageSize = (128,128))\n####################","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:32:16.151850Z","iopub.execute_input":"2021-06-23T19:32:16.152193Z","iopub.status.idle":"2021-06-23T19:56:33.270800Z","shell.execute_reply.started":"2021-06-23T19:32:16.152158Z","shell.execute_reply":"2021-06-23T19:56:33.269078Z"},"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Getting Numpy Array From Test Data Set","metadata":{}},{"cell_type":"code","source":"X = np.load('ImageArray128.npy')\ny = np.load('TargetArray128.npy')","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:02:26.581012Z","iopub.execute_input":"2021-06-23T20:02:26.581354Z","iopub.status.idle":"2021-06-23T20:02:26.621373Z","shell.execute_reply.started":"2021-06-23T20:02:26.581304Z","shell.execute_reply":"2021-06-23T20:02:26.620497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"directory = '../input/siim-covid19-detection/test'\n# getting all the image paths\nall_abs_file_paths = []\nfor dirpath,_,filenames in os.walk(directory):\n        for f in filenames:\n            all_abs_file_paths.append(os.path.abspath(os.path.join(dirpath, f)))\n\ntest_df = pd.DataFrame(all_abs_file_paths, columns = ['abs_image_path'])\ntest_df['id'] = test_df['abs_image_path'].apply(lambda x: x.split(\"/\")[5])\ntest_df['abs_image_path'] = test_df['abs_image_path'].apply(lambda x: [x])\nX_final_test = getNumpyArraysSet(test_df) # test data by Kaggle","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:02:29.960816Z","iopub.execute_input":"2021-06-23T20:02:29.961151Z","iopub.status.idle":"2021-06-23T20:02:39.804910Z","shell.execute_reply.started":"2021-06-23T20:02:29.961120Z","shell.execute_reply":"2021-06-23T20:02:39.802567Z"},"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Preparation For Training the Model\n- Once you have have numpy arrays you canuse any model or train the data, We will be using efficientNet to train our first model. \n- Hoever before we move forward we must prepare the numpy arrays for the pretrained model.","metadata":{"execution":{"iopub.status.busy":"2021-06-23T19:09:51.440796Z","iopub.execute_input":"2021-06-23T19:09:51.441305Z","iopub.status.idle":"2021-06-23T19:09:51.447552Z","shell.execute_reply.started":"2021-06-23T19:09:51.441200Z","shell.execute_reply":"2021-06-23T19:09:51.446125Z"}}},{"cell_type":"markdown","source":"## Minimal data augmentation for upsampling the minority class","metadata":{"execution":{"iopub.status.busy":"2021-06-22T00:12:12.513606Z","iopub.execute_input":"2021-06-22T00:12:12.513967Z","iopub.status.idle":"2021-06-22T00:17:05.716171Z","shell.execute_reply.started":"2021-06-22T00:12:12.513934Z","shell.execute_reply":"2021-06-22T00:17:05.715112Z"}}},{"cell_type":"code","source":"transform_train = transforms.Compose([transforms.ToPILImage(),\n                    transforms.RandomApply([torchvision.transforms.RandomRotation(10),transforms.RandomHorizontalFlip()],0.7), \n                    transforms.ToTensor()])","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:02:47.073530Z","iopub.execute_input":"2021-06-23T20:02:47.073884Z","iopub.status.idle":"2021-06-23T20:02:47.080663Z","shell.execute_reply.started":"2021-06-23T20:02:47.073852Z","shell.execute_reply":"2021-06-23T20:02:47.079687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#upsampling the minority class\nnewX = []\nnewY = []\nfor j in range(X.shape[-1]):\n    if y[0,-1,j] == 1:\n        for k in range(8):\n            newX.append(transform_train(X[:,:,j]))\n            newY.append(y[0,:,j])\n    else:\n        newX.append(transform_train(X[:,:,j]))\n        newY.append(y[0,:,j])\nbalancedX = np.stack(newX)\nbalancedX = np.repeat(balancedX, repeats = 3, axis =1)\nbalancedY = np.stack(newY)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:02:49.171779Z","iopub.execute_input":"2021-06-23T20:02:49.172104Z","iopub.status.idle":"2021-06-23T20:02:55.741710Z","shell.execute_reply.started":"2021-06-23T20:02:49.172073Z","shell.execute_reply":"2021-06-23T20:02:55.740823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# checking unique values in new dataset.\ntarget_df = pd.DataFrame(balancedY, columns = ['Negative for Pneumonia', 'Typical Appearance',\n           'Indeterminate Appearance', 'Atypical Appearance'])\ntarget_df.value_counts()","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:02:55.743452Z","iopub.execute_input":"2021-06-23T20:02:55.743790Z","iopub.status.idle":"2021-06-23T20:02:55.759460Z","shell.execute_reply.started":"2021-06-23T20:02:55.743755Z","shell.execute_reply":"2021-06-23T20:02:55.758712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Splitting into training, test and Validation set","metadata":{}},{"cell_type":"code","source":"## Splitting into training and validation set.\nvalidFrac = 0.1   # Define the fraction of images to move to validation dataset\ntestFrac = 0.1    # Define the fraction of images to move to test dataset\nvalidList = []\ntestList = []\ntrainList = []\n\nfor i in range(len(balancedX)):\n    rann = np.random.random() # Randomly reassign images\n    if rann < validFrac:\n        validList.append(i)\n    elif rann < testFrac + validFrac:\n        testList.append(i)\n    else:\n        trainList.append(i)\n        \nnTrain = len(trainList)  # Count the number in each set\nnValid = len(validList)\nnTest = len(testList)\nprint(\"Training images =\",nTrain,\"Validation =\",nValid,\"Testing =\",nTest)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:02:59.673052Z","iopub.execute_input":"2021-06-23T20:02:59.673390Z","iopub.status.idle":"2021-06-23T20:02:59.701765Z","shell.execute_reply.started":"2021-06-23T20:02:59.673358Z","shell.execute_reply":"2021-06-23T20:02:59.700292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Converting numpy arrays into torch tensors......... Data that is understood by tensorflow and pytorch","metadata":{}},{"cell_type":"code","source":"trainIds = torch.tensor(trainList)    # Slice the big image and label tensors up into\nvalidIds = torch.tensor(validList) \ntestIds = torch.tensor(testList)       #       training, validation, and testing tensors\ntrainX = torch.tensor(balancedX[trainIds,:,:,:])\ntrainY = torch.tensor(balancedY[trainIds,:])\nvalidX = torch.tensor(balancedX[validIds,:,:,:])\nvalidY = torch.tensor(balancedY[validIds,:])\ntestX = torch.tensor(balancedX[testIds,:,:,:])\ntestY = torch.tensor(balancedY[testIds,:])","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:04:14.306515Z","iopub.execute_input":"2021-06-23T20:04:14.306853Z","iopub.status.idle":"2021-06-23T20:04:15.776079Z","shell.execute_reply.started":"2021-06-23T20:04:14.306825Z","shell.execute_reply":"2021-06-23T20:04:15.775167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading the EfficientNet Model","metadata":{}},{"cell_type":"code","source":"model = EfficientNet.from_pretrained('efficientnet-b2', num_classes=4) # loading model\nmodel.to(dev) # Sendeng the model to GPU","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:04:19.577470Z","iopub.execute_input":"2021-06-23T20:04:19.577786Z","iopub.status.idle":"2021-06-23T20:04:28.086747Z","shell.execute_reply.started":"2021-06-23T20:04:19.577756Z","shell.execute_reply":"2021-06-23T20:04:28.085891Z"},"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Explore the model complexity\nprint(summary(model, input_size=(3, 128, 128)))","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:04:28.088262Z","iopub.execute_input":"2021-06-23T20:04:28.088657Z","iopub.status.idle":"2021-06-23T20:04:29.062803Z","shell.execute_reply.started":"2021-06-23T20:04:28.088616Z","shell.execute_reply":"2021-06-23T20:04:29.061978Z"},"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training the Model","metadata":{}},{"cell_type":"code","source":"numClass = len(balancedY[:,0])\nlearnRate = 0.01          # Define a learning rate.\nmaxEpochs = 20            # Maximum training epochs\nt2vRatio = 1.5            # Maximum allowed ratio of validation to training loss\nt2vEpochs = 3            # Number of consecutive epochs before halting if validation loss exceeds above limit\nbatchSize = 128           # Batch size. Going too large will cause an out-of-memory error.\ntrainBats = nTrain // batchSize       # Number of training batches per epoch. Round down to simplify last batch\nvalidBats = nValid // batchSize       # Validation batches. Round down\ntestBats = -(-nTest // batchSize)     # Testing batches. Round up to include all\nopti = om.SGD(model.parameters(), lr = learnRate)   # Initialize an optimizer\nlossEval = nn.CrossEntropyLoss()\n\nfor i in range(maxEpochs):\n    model.train()                     # Set model to training mode\n    epochLoss = 0.\n    permute = torch.randperm(nTrain)  # Shuffle data to randomize batches\n    trainX = trainX[permute,:,:,:]\n    trainY = trainY[permute]\n    for j in range(trainBats):        # Iterate over batches\n        if j%20 == 0:\n          print('The batch num is {}'.format(j))\n        opti.zero_grad()              # Zero out gradient accumulated in optimizer\n        batX = trainX[j*batchSize:(j+1)*batchSize,:,:,:].to(dev)   # Slice shuffled data into batches\n        batY = trainY[j*batchSize:(j+1)*batchSize].to(dev)         # .to(dev) moves these batches to the GPU\n        yOut = model(batX)            # Evalute predictions\n        # print(yOut, batY)\n        loss = F.cross_entropy(yOut, torch.max(batY, 1)[1])        # Compute loss\n        epochLoss += loss.item()      # Add loss\n        loss.backward()               # Backpropagate loss\n        opti.step()                   # Update model weights using optimizer\n    validLoss = 0.\n    permute = torch.randperm(nValid)  # We go through the exact same steps, without backprop / optimization\n    validX = validX[permute,:,:,:]    # in order to evaluate the validation loss\n    validY = validY[permute]\n    model.eval()                      # Set model to evaluation mode\n    with torch.no_grad():             # Temporarily turn off gradient descent\n        for j in range(validBats):\n            opti.zero_grad()\n            batX = validX[j*batchSize:(j+1)*batchSize,:,:,:].to(dev)\n            batY = validY[j*batchSize:(j+1)*batchSize].to(dev)\n            yOut = model(batX)\n            validLoss += F.cross_entropy(yOut, torch.max(batY, 1)[1]).item()\n    epochLoss /= trainBats            # Average loss over batches and print\n    validLoss /= validBats\n    print(\"Epoch = {:-3}; Training loss = {:.4f}; Validation loss = {:.4f}\".format(i,epochLoss,validLoss))\n    if validLoss > t2vRatio * epochLoss:\n        t2vEpochs -= 1                # Test if validation loss exceeds halting threshold\n        if t2vEpochs < 1:\n            print(\"Validation loss too high; halting to prevent overfitting\")\n            break","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:04:47.153600Z","iopub.execute_input":"2021-06-23T20:04:47.153923Z","iopub.status.idle":"2021-06-23T20:11:01.936894Z","shell.execute_reply.started":"2021-06-23T20:04:47.153893Z","shell.execute_reply":"2021-06-23T20:11:01.936031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Testing the model on Test Data","metadata":{}},{"cell_type":"code","source":"confuseMtx = np.zeros((numClass,numClass),dtype=int)    # Create empty confusion matrix\nmodel.eval()\npredictions = []\nwith torch.no_grad():\n    permute = torch.randperm(nTest)                     # Shuffle test data\n    testX = testX[permute,:,:,:]\n    testY = testY[permute]\n    for j in range(testBats):                           # Iterate over test batches\n        batX = testX[j*batchSize:(j+1)*batchSize,:,:,:].to(dev)\n        batY = testY[j*batchSize:(j+1)*batchSize].to(dev)\n        yOut = model(batX) # Pass test batch through model\n        for j in yOut:\n            print(j, len(j),np.zeros(len(j)), torch.max(j,0)[1])\n            b = np.zeros(len(j))\n            index = torch.max(j,0)[1]\n            b[index] = 1\n            predictions.append(b)\n        pred = yOut.max(1,keepdim=True)[1]              # Generate predictions by finding the max Y values\n        for j in torch.cat((batY.max(1,keepdim=True)[1], pred),dim=1).tolist(): # Glue together Actual and Predicted to\n            confuseMtx[j[0],j[1]] += 1                  # make (row, col) pairs, and increment confusion matrix\ncorrect = sum([confuseMtx[i,i] for i in range(numClass)])   # Sum over diagonal elements to count correct predictions\nprint(\"Correct predictions: \",correct,\"of\",nTest)\nprint(\"Confusion Matrix:\")\nprint(confuseMtx)\n# print(classNames)","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:26:02.538083Z","iopub.execute_input":"2021-06-23T20:26:02.538435Z","iopub.status.idle":"2021-06-23T20:26:05.158262Z","shell.execute_reply.started":"2021-06-23T20:26:02.538401Z","shell.execute_reply":"2021-06-23T20:26:05.157485Z"},"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testX","metadata":{"execution":{"iopub.status.busy":"2021-06-23T21:36:50.149835Z","iopub.execute_input":"2021-06-23T21:36:50.150538Z","iopub.status.idle":"2021-06-23T21:36:50.397347Z","shell.execute_reply.started":"2021-06-23T21:36:50.150443Z","shell.execute_reply":"2021-06-23T21:36:50.395996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get confusion matrix using scikit-learn\nfrom sklearn.metrics import confusion_matrix\npredictions = np.stack(predictions)\nconfusion_matrix()","metadata":{"execution":{"iopub.status.busy":"2021-06-23T20:26:27.658244Z","iopub.execute_input":"2021-06-23T20:26:27.658602Z","iopub.status.idle":"2021-06-23T20:26:27.774832Z","shell.execute_reply.started":"2021-06-23T20:26:27.658571Z","shell.execute_reply":"2021-06-23T20:26:27.774060Z"},"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]}]}