{"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":"# CNN Image Classification\nThe model in this notebook will be trained to classify between two types of images: One with malignent cancer cells, and one with normal cancer cells.\nThe pathology images used for training and testing can be found on [Kaggle](https://www.kaggle.com/competitions/histopathologic-cancer-detection/data) ","metadata":{}},{"cell_type":"markdown","source":"<a name='1'></a>\n## Part 1: Data Preprocessing","metadata":{}},{"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 matplotlib.pyplot as plt\nimport torch\nfrom plotly.subplots import make_subplots\nimport plotly.graph_objs as go\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score\nimport copy\nimport os\nimport torch\nfrom PIL import Image\nfrom PIL import Image, ImageDraw\nfrom torch.utils.data import Dataset\nimport torchvision.transforms as transforms\nfrom torch.utils.data import random_split\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport torch.nn as nn\nimport torchvision\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n#device = torch.device('cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:14:36.981424Z","iopub.execute_input":"2023-02-05T00:14:36.982168Z","iopub.status.idle":"2023-02-05T00:14:40.141360Z","shell.execute_reply.started":"2023-02-05T00:14:36.982077Z","shell.execute_reply":"2023-02-05T00:14:40.140182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a name='1.1'></a>\n### Part 1.1: EDA","metadata":{}},{"cell_type":"markdown","source":"Note: A positive label indicates that the center 32x32px region of a patch contains at least one pixel of tumor tissue. Tumor tissue in the outer region of the patch does not influence the label. This outer region is provided to enable fully-convolutional models that do not use zero-padding, to ensure consistent behavior when applied to a whole-slide image.","metadata":{}},{"cell_type":"code","source":"# Checking the attributes and length of labels dataframe\ntrain_img_location = \"/kaggle/input/histopathologic-cancer-detection/train\"\nlabels = pd.read_csv('/kaggle/input/histopathologic-cancer-detection/train_labels.csv')\nprint(len(labels))\nlabels.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:14:40.143581Z","iopub.execute_input":"2023-02-05T00:14:40.144501Z","iopub.status.idle":"2023-02-05T00:14:40.777309Z","shell.execute_reply.started":"2023-02-05T00:14:40.144458Z","shell.execute_reply":"2023-02-05T00:14:40.776258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Note to self: It is important to check the class distribution in the training set to avoid bad generalization","metadata":{}},{"cell_type":"code","source":"labels['label'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:14:40.778970Z","iopub.execute_input":"2023-02-05T00:14:40.779503Z","iopub.status.idle":"2023-02-05T00:14:40.795717Z","shell.execute_reply.started":"2023-02-05T00:14:40.779469Z","shell.execute_reply":"2023-02-05T00:14:40.794358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Visualization of the training dataset images","metadata":{}},{"cell_type":"code","source":"malignant = labels.loc[labels['label']==1]['id'].to_numpy()    # get the ids of malignant cases\nnormal = labels.loc[labels['label']==0]['id'].to_numpy()       # get the ids of the normal cases\n\n# Create a subplot to display train examples\nnrows,ncols=6,15\nfig,ax = plt.subplots(nrows,ncols,figsize=(15,6))\nplt.subplots_adjust(wspace=0, hspace=0) \n\n# Get the last nrows*ncols images and display\nfor i,j in enumerate(malignant[:nrows*ncols]):\n    fname = os.path.join(train_img_location ,j +'.tif')\n    img = Image.open(fname)\n    idcol = ImageDraw.Draw(img)\n    idcol.rectangle(((0,0),(95,95)),outline='red')\n    plt.subplot(nrows, ncols, i+1) \n    plt.imshow(np.array(img))\n    plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:14:40.800005Z","iopub.execute_input":"2023-02-05T00:14:40.800358Z","iopub.status.idle":"2023-02-05T00:14:45.224184Z","shell.execute_reply.started":"2023-02-05T00:14:40.800328Z","shell.execute_reply":"2023-02-05T00:14:45.223148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"nrows,ncols=6,15\nfig,ax = plt.subplots(nrows,ncols,figsize=(15,6))\nplt.subplots_adjust(wspace=0, hspace=0) \n\nfor i,j in enumerate(normal[:nrows*ncols]):\n    fname = os.path.join(train_img_location ,j +'.tif')\n    img = Image.open(fname)\n    idcol = ImageDraw.Draw(img)\n    idcol.rectangle(((0,0),(95,95)),outline='green')\n    plt.subplot(nrows, ncols, i+1) \n    plt.imshow(np.array(img))\n    plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:14:45.226074Z","iopub.execute_input":"2023-02-05T00:14:45.226496Z","iopub.status.idle":"2023-02-05T00:14:49.780409Z","shell.execute_reply.started":"2023-02-05T00:14:45.226458Z","shell.execute_reply":"2023-02-05T00:14:49.779396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Part 1.2: Prepare dataset","metadata":{}},{"cell_type":"code","source":"torch.manual_seed(0)\nclass HistopathologicCancerDataset(Dataset):\n    \"\"\"\n    This is our custom dataset class which will load the images, perform transforms on them,\n    and load their corresponding labels.\n    To create this, we need to override the __init__, __getitem__, and __len__ methods.\n    Variables:\n        labels_df(DataFrame): Has 2 attribute: Image_name, Label\n        img_dir(String): Path to the image directory\n        images(List): Full name of the images to be used\n    \"\"\"\n    def __init__(self, img_dir, labels_csv_file, max_size=None, transform=None):\n        self.img_dir = img_dir\n        self.transform = transform\n        self.images = []\n        labels_df = pd.read_csv(labels_csv_file)\n        self.labels_df = labels_df.sample(frac=1).reset_index(drop=True)\n\n        for i,item in self.labels_df.iterrows():\n            f = item['id'] + \".tif\"\n            if f in os.listdir(img_dir): # If the image can be found in img_dir\n                self.images.append(item['id']) # Add its image id to self.images\n            else:\n                self.labels_df.drop(i,axis=0,inplace=True) # Else delete the row from self.labels_df\n            if max_size and len(self.images) == max_size:\n                break            \n        \n        #self.labels_df.reset_index(drop=True,inplace=True) # Sync up the index of self.labels_df and self.images\n\n    def __getitem__(self, idx):\n        \"\"\"\n        Open image, apply transforms and return with label\n        Return: Object sample with 3 attributes:\n            image: Image tensor\n            label: Image label(0 or 1)\n            id: Image name(id)\n        \"\"\"\n        img_name = self.images[idx]\n        img_path = os.path.join(self.img_dir, img_name + \".tif\")\n        image = Image.open(img_path)  # Open Image with PIL\n        image = self.transform(image) # Apply Specific Transformation to Image\n        correct_label = self.labels_df.loc[self.labels_df['id'] == img_name]['label'].item()\n\n        return image, correct_label\n    \n    def __len__(self):\n        return len(self.images)","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:14:49.782118Z","iopub.execute_input":"2023-02-05T00:14:49.782494Z","iopub.status.idle":"2023-02-05T00:14:49.798519Z","shell.execute_reply.started":"2023-02-05T00:14:49.782459Z","shell.execute_reply":"2023-02-05T00:14:49.797384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Create train and validation dataloader","metadata":{}},{"cell_type":"code","source":"img_dataset = HistopathologicCancerDataset(\n    img_dir=train_img_location,\n    labels_csv_file=\"/kaggle/input/histopathologic-cancer-detection/train_labels.csv\",\n    transform=transforms.ToTensor(),\n    max_size=4000\n)\n# Get a sample from the dataset\nsample = img_dataset[10]\nprint(\"Size of item 0 from img_dataset:\")\nprint(sample[0].shape, sample[1])","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:14:49.800369Z","iopub.execute_input":"2023-02-05T00:14:49.801156Z","iopub.status.idle":"2023-02-05T00:26:30.619584Z","shell.execute_reply.started":"2023-02-05T00:14:49.801103Z","shell.execute_reply":"2023-02-05T00:26:30.618435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here we split dataset into train + valid dataset and put them into a DataLoader","metadata":{}},{"cell_type":"code","source":"len_img=len(img_dataset)\nlen_train=int(0.8*len_img)\nlen_val=len_img-len_train\n\n# Split Pytorch tensor\ntrain_ts,val_ts=random_split(img_dataset, [len_train,len_val])\n\nprint(\"train dataset size:\", len(train_ts))\nprint(\"validation dataset size:\", len(val_ts))","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:30.621025Z","iopub.execute_input":"2023-02-05T00:26:30.621389Z","iopub.status.idle":"2023-02-05T00:26:30.629238Z","shell.execute_reply.started":"2023-02-05T00:26:30.621359Z","shell.execute_reply":"2023-02-05T00:26:30.628232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# getting the torch tensor image & target variable\nii=-1\nfor x,y in train_ts:\n    print(x.shape,y)\n    ii+=1\n    if(ii>5):\n        break","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:30.630958Z","iopub.execute_input":"2023-02-05T00:26:30.631348Z","iopub.status.idle":"2023-02-05T00:26:31.031144Z","shell.execute_reply.started":"2023-02-05T00:26:30.631317Z","shell.execute_reply":"2023-02-05T00:26:31.029983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the following transformations for the training dataset\ntr_transf = transforms.Compose([\n#     transforms.Resize((40,40)),\n    # transforms.RandomHorizontalFlip(p=0.5), \n    # transforms.RandomVerticalFlip(p=0.5),  \n    # transforms.RandomRotation(45),         \n#     transforms.RandomResizedCrop(50,scale=(0.8,1.0),ratio=(1.0,1.0)),\n    transforms.ToTensor()])\n\n# For the validation dataset, we don't need any augmentation; simply convert images into tensors\nval_transf = transforms.Compose([\n    transforms.ToTensor()])\n\n# After defining the transformations, overwrite the transform functions of train_ts, val_ts\ntrain_ts.transform=tr_transf\nval_ts.transform=val_transf","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:31.034271Z","iopub.execute_input":"2023-02-05T00:26:31.034719Z","iopub.status.idle":"2023-02-05T00:26:31.040880Z","shell.execute_reply.started":"2023-02-05T00:26:31.034686Z","shell.execute_reply":"2023-02-05T00:26:31.039813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from random import shuffle\ntrain_loader = torch.utils.data.DataLoader(\n    train_ts,\n    batch_size=64,\n    shuffle = True,\n)\nval_loader = torch.utils.data.DataLoader(\n    val_ts,\n    batch_size=64,\n    shuffle = False,\n)","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:31.042119Z","iopub.execute_input":"2023-02-05T00:26:31.042703Z","iopub.status.idle":"2023-02-05T00:26:31.056372Z","shell.execute_reply.started":"2023-02-05T00:26:31.042665Z","shell.execute_reply":"2023-02-05T00:26:31.055144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Check the dimension of objects in the dataset/dataloader","metadata":{}},{"cell_type":"code","source":"for i in train_loader:\n    print(i) # A batch of image matrices\n    #print(i[0].shape) # A batch of correct label\n    #print(len(i['id'])) # A batch of image id\n    break\nprint(train_ts[0]) # Show an example in the dataset train_ts","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:31.057822Z","iopub.execute_input":"2023-02-05T00:26:31.058203Z","iopub.status.idle":"2023-02-05T00:26:34.634540Z","shell.execute_reply.started":"2023-02-05T00:26:31.058159Z","shell.execute_reply":"2023-02-05T00:26:34.633364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a name='2'></a>\n## Part 2: Model definition","metadata":{}},{"cell_type":"markdown","source":"Here we use the model ResNet50.","metadata":{}},{"cell_type":"code","source":"model = torchvision.models.resnet34().to(device)\nmodel","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:34.636079Z","iopub.execute_input":"2023-02-05T00:26:34.636538Z","iopub.status.idle":"2023-02-05T00:26:35.018981Z","shell.execute_reply.started":"2023-02-05T00:26:34.636495Z","shell.execute_reply":"2023-02-05T00:26:35.017838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.fc = torch.nn.Sequential(\n    torch.nn.Linear(\n        in_features=512,\n        out_features=1\n    ),\n    torch.nn.Sigmoid(),\n)\nmodel","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:35.020430Z","iopub.execute_input":"2023-02-05T00:26:35.021463Z","iopub.status.idle":"2023-02-05T00:26:35.031095Z","shell.execute_reply.started":"2023-02-05T00:26:35.021424Z","shell.execute_reply":"2023-02-05T00:26:35.029981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Functions for training and evelauating\n# Function to get the learning rate\ndef get_lr(opt):\n    for param_group in opt.param_groups:\n        return param_group['lr']\n\n# Function to compute the loss value per batch of data\ndef loss_batch(loss_func, output, target, opt=None):\n    loss = loss_func(output, target.view(-1, 1).float()) # get loss\n    pred = (output>=0.5) # Get Output Class\n    metric_b=torch.sum((pred == target.view(-1, 1))) # get performance metric\n    \n    if opt is not None:\n        opt.zero_grad()\n        loss.backward()\n        opt.step()\n\n    return loss.item(), metric_b\n\n# Compute the loss value & performance metric for the entire dataset (epoch)\ndef loss_epoch(model,loss_func,dataset_dl,check=False,opt=None):\n    \n    run_loss=0.0 \n    t_metric=0.0\n    len_data=len(dataset_dl.dataset)\n\n    # internal loop over dataset\n    for xb, yb in dataset_dl:\n        # move batch to device\n        xb=xb.to(device)\n        yb=yb.to(device)\n        output=model(xb) # get model output\n        loss_b,metric_b=loss_batch(loss_func, output, yb, opt) # get loss per batch\n        run_loss+=loss_b        # update running loss\n\n        if metric_b is not None: # update running metric\n            t_metric+=metric_b\n\n        # break the loop in case of sanity check\n        if check is True:\n            break\n    \n    loss=run_loss/float(len_data)  # average loss value\n    metric=t_metric/float(len_data) # average metric value\n    \n    return loss, metric","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:35.033071Z","iopub.execute_input":"2023-02-05T00:26:35.033512Z","iopub.status.idle":"2023-02-05T00:26:35.046349Z","shell.execute_reply.started":"2023-02-05T00:26:35.033470Z","shell.execute_reply":"2023-02-05T00:26:35.045248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_val(model, params,verbose=False):\n    \n    # Get the parameters\n    epochs=params[\"epochs\"]\n    loss_func=params[\"f_loss\"]\n    opt=params[\"optimiser\"]\n    train_dl=params[\"train\"]\n    val_dl=params[\"val\"]\n    check=params[\"check\"]\n    lr_scheduler=params[\"lr_change\"]\n    weight_path=params[\"weight_path\"]\n    \n    loss_history={\"train\": [],\"val\": []} # history of loss values in each epoch\n    metric_history={\"train\": [],\"val\": []} # histroy of metric values in each epoch\n    best_model_wts = copy.deepcopy(model.state_dict()) # a deep copy of weights for the best performing model\n    best_loss=float('inf') # initialize best loss to a large value\n    \n    # main loop\n    for epoch in range(epochs):\n        \n        ''' Get the Learning Rate '''\n        current_lr=get_lr(opt)\n        if(verbose):\n            print('Epoch {}/{}, current lr={}'.format(epoch, epochs - 1, current_lr))\n        \n        ''' Train the Model on the Training Set '''\n        model.train()\n        train_loss, train_metric=loss_epoch(model,loss_func,train_dl,check,opt)\n\n        ''' Collect loss and metric for training dataset ''' \n        loss_history[\"train\"].append(train_loss)\n        metric_history[\"train\"].append(train_metric)\n        \n        ''' Evaluate model on validation dataset '''\n        model.eval()\n        with torch.no_grad():\n            val_loss, val_metric=loss_epoch(model,loss_func,val_dl,check)\n        \n        # store best model\n        if val_loss < best_loss:\n            best_loss = val_loss\n            best_model_wts = copy.deepcopy(model.state_dict())\n            \n            # store weights into a local file\n            torch.save(model.state_dict(), weight_path)\n            if(verbose):\n                print(\"Copied best model weights!\")\n        \n        # collect loss and metric for validation dataset\n        loss_history[\"val\"].append(val_loss)\n        metric_history[\"val\"].append(val_metric)\n        \n        # learning rate schedule\n        lr_scheduler.step(val_loss)\n        if current_lr != get_lr(opt):\n            if(verbose):\n                print(\"Loading best model weights!\")\n            model.load_state_dict(best_model_wts) \n\n        if(verbose):\n            print(f\"train loss: {train_loss:.6f}, dev loss: {val_loss:.6f}, accuracy: {100*val_metric:.3f}\")\n            print(\"-\"*10) \n\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n        \n    return model, loss_history, metric_history\n\nparams_train={\n \"train\": train_loader,\"val\": val_loader,\n \"epochs\": 50,\n \"optimiser\": torch.optim.Adam(model.parameters(),\n                         lr=3e-4),\n \"lr_change\": ReduceLROnPlateau(torch.optim.Adam(model.parameters(), lr=3e-4),\n                                mode='min',\n                                factor=0.5,\n                                patience=20,\n                                verbose=0),\n \"f_loss\": nn.BCELoss(),\n \"weight_path\": \"resnet34.pt\",\n \"check\": False, \n}\n\n''' Actual Train / Evaluation of CNN Model '''\n# train and validate the model\ncnn_model,loss_hist,metric_hist=train_val(model.to(device),params_train,verbose=True)","metadata":{"execution":{"iopub.status.busy":"2023-02-05T00:26:35.048737Z","iopub.execute_input":"2023-02-05T00:26:35.050107Z","iopub.status.idle":"2023-02-05T05:08:32.216150Z","shell.execute_reply.started":"2023-02-05T00:26:35.050066Z","shell.execute_reply":"2023-02-05T05:08:32.214983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc \ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-02-05T05:08:32.217614Z","iopub.execute_input":"2023-02-05T05:08:32.218000Z","iopub.status.idle":"2023-02-05T05:08:32.556436Z","shell.execute_reply.started":"2023-02-05T05:08:32.217962Z","shell.execute_reply":"2023-02-05T05:08:32.555193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train-Validation Progress\nepochs=params_train[\"epochs\"]\nmetric_hist[\"train\"] = [i.cpu() for i in metric_hist[\"train\"]]\nmetric_hist[\"val\"] = [i.cpu() for i in metric_hist[\"val\"]]\n\nfig = make_subplots(rows=1, cols=2,subplot_titles=['Loss history','Accuracy history'])\nfig.add_trace(go.Scatter(x=[*range(1,epochs+1)], y=loss_hist[\"train\"],name='Training'),row=1, col=1)\nfig.add_trace(go.Scatter(x=[*range(1,epochs+1)], y=loss_hist[\"val\"],name='Validating'),row=1, col=1)\nfig.add_trace(go.Scatter(x=[*range(1,epochs+1)], y=metric_hist[\"train\"],name='Training'),row=1, col=2)\nfig.add_trace(go.Scatter(x=[*range(1,epochs+1)], y=metric_hist[\"val\"],name='Validating'),row=1, col=2)\nfig.update_layout(template='plotly_white');fig.update_layout(margin={\"r\":0,\"t\":60,\"l\":0,\"b\":0},height=300)\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-05T05:08:32.557968Z","iopub.execute_input":"2023-02-05T05:08:32.558379Z","iopub.status.idle":"2023-02-05T05:08:33.693149Z","shell.execute_reply.started":"2023-02-05T05:08:32.558345Z","shell.execute_reply":"2023-02-05T05:08:33.691656Z"},"trusted":true},"execution_count":null,"outputs":[]}]}