{"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":"#@title Import Statements (Run This Cell)\n\nimport os\nimport pandas as pd\nimport numpy as np\nfrom sklearn.model_selection import train_test_split\nfrom pydicom import dcmread\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom sklearn.decomposition import PCA\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.preprocessing import StandardScaler\n\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.transforms as transforms\nimport torch.optim as optim\nfrom tqdm.notebook import tqdm, trange\nimport torch.nn.functional as F\nfrom torch.utils import data","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nPlotting Functions\n'''\n\ndef plot_loss_accuracy(train_loss, train_acc, validation_loss, validation_acc):\n    epochs = [i+1 for i in range(len(train_loss))]\n    fig, (ax1, ax2) = plt.subplots(1, 2)\n    ax1.plot(epochs, train_loss, label='Training Loss')\n    ax1.plot(epochs, validation_loss, label='Validation Loss')\n    ax1.set_xlabel('Epochs')\n    ax1.set_ylabel('Loss')\n    ax1.set_title('Epoch vs Loss')\n    ax1.set_xticks(epochs)\n    ax1.legend()\n\n    ax2.plot(epochs, train_acc, label='Training Accuracy')\n    ax2.plot(epochs, validation_acc, label='Validation Accuracy')\n    ax2.set_xlabel('Epochs')\n    ax2.set_ylabel('Accuracy')\n    ax2.set_title('Epoch vs Accuracy')\n    ax2.set_xticks(epochs)\n    ax2.legend()\n    fig.set_size_inches(15.5, 5.5)\n    plt.show()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preparing labels","metadata":{}},{"cell_type":"code","source":"label_data = pd.read_csv('../input/rsna-pneumonia-detection-challenge/stage_2_train_labels.csv')\ncolumns = ['patientId', 'Target']\n\nlabel_data = label_data.filter(columns)\nlabel_data.head(5)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dividing labels for train and validation set","metadata":{}},{"cell_type":"code","source":"train_labels, val_labels = train_test_split(label_data.values, test_size=0.1)\nprint(train_labels.shape)\nprint(val_labels.shape)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'patientId: {train_labels[0][0]}, Target: {train_labels[0][1]}')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preparing train and validation image paths","metadata":{}},{"cell_type":"code","source":"train_f = '../input/rsna-pneumonia-detection-challenge/stage_2_train_images'\ntest_f = '../input/rsna-pneumonia-detection-challenge/stage_2_test_images'\n\ntrain_paths = [os.path.join(train_f, image[0]) for image in train_labels]\nval_paths = [os.path.join(train_f, image[0]) for image in val_labels]\n\nprint(len(train_paths))\nprint(len(val_paths))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Show some samples from data","metadata":{}},{"cell_type":"code","source":"def imshow(num_to_show=9):\n    \n    plt.figure(figsize=(10,10))\n    \n    for i in range(num_to_show):\n        plt.subplot(3, 3, i+1)\n        plt.grid(False)\n        plt.xticks([])\n        plt.yticks([])\n        \n        img_dcm = dcmread(f'{train_paths[i+20]}.dcm')\n        img_np = img_dcm.pixel_array\n        plt.imshow(img_np, cmap=plt.cm.binary)\n        plt.xlabel(train_labels[i+20][1])\n\nimshow()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Composing transformations","metadata":{}},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.RandomHorizontalFlip(),\n    transforms.Resize(224),\n    transforms.ToTensor()])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Write a custom dataset ","metadata":{}},{"cell_type":"code","source":"class Dataset(data.Dataset):\n    \n    def __init__(self, paths, labels, transform=None):\n        self.paths = paths\n        self.labels = labels\n        self.transform = transform\n    \n    def __getitem__(self, index):\n        image = dcmread(f'{self.paths[index]}.dcm')\n        image = image.pixel_array\n        image = image / 255.0\n\n        image = (255*image).clip(0, 255).astype(np.uint8)\n        image = Image.fromarray(image).convert('RGB')\n\n        label = self.labels[index][1]\n        \n        if self.transform is not None:\n            image = self.transform(image)\n            \n        return image, label\n    \n    def __len__(self):\n        \n        return len(self.paths)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check the custom dataset","metadata":{}},{"cell_type":"code","source":"train_dataset = Dataset(train_paths, train_labels, transform=transform)\nimage = iter(train_dataset)\nimg, label = next(image)\nprint(f'Tensor:{img}, Label:{label}')\nimg = np.transpose(img, (1, 2, 0))\nplt.imshow(img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train image shape","metadata":{}},{"cell_type":"code","source":"img.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prepare training and validation dataloader","metadata":{}},{"cell_type":"code","source":"train_dataset = Dataset(train_paths, train_labels, transform=transform)\nval_dataset = Dataset(val_paths, val_labels, transform=transform)\ntrain_loader = data.DataLoader(dataset=train_dataset, batch_size=128, shuffle=True)\nval_loader = data.DataLoader(dataset=val_dataset, batch_size=128, shuffle=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check dataloader","metadata":{}},{"cell_type":"code","source":"batch = iter(train_loader)\nimages, labels = next(batch)\n\nimage_grid = torchvision.utils.make_grid(images[:4])\nimage_np = image_grid.numpy()\nimg = np.transpose(image_np, (1, 2, 0))\nplt.imshow(img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Specify device object","metadata":{}},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining the CNN","metadata":{}},{"cell_type":"code","source":"class Baseline_CNN(nn.Module):\n    def __init__(self):\n        super(Baseline_CNN, self).__init__()\n        self.conv1 = nn.Conv2d(3, 32, 3, 1)\n        self.conv2 = nn.Conv2d(32, 64, 3, 1)\n        self.fc1 = nn.Linear(774400, 128)\n        self.fc2 = nn.Linear(128, 2)\n        \n    def forward(self, x):\n        x = self.conv1(x)\n        x = F.relu(x)\n        x = self.conv2(x)\n        x = F.relu(x)\n        x = F.max_pool2d(x, 2)\n        # x = F.max_pool2d(x, 4)\n        x = torch.flatten(x, 1)\n        x = self.fc1(x)\n        x = F.relu(x)\n        x = self.fc2(x)  \n        return x","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(Baseline_CNN)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training Function\nThis function was taken from the CIS 522 Lectures","metadata":{}},{"cell_type":"code","source":"def train(model, device, train_loader, validation_loader, epochs):\n\n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.SGD(model.parameters(),\n                              lr=0.01, momentum=0.9)\n    \n    train_loss, validation_loss = [], []\n    train_acc, validation_acc = [], []\n    for epoch in range(epochs):\n        model.train()\n        running_loss = 0.\n        correct, total = 0, 0 \n        with tqdm(train_loader, unit='batch') as tepoch:\n            tepoch.set_description('Training: ')\n            for data, target in tepoch:\n                data, target = data.to(device), target.to(device)\n\n                # add micro for coding training loop\n                optimizer.zero_grad()\n                output = model(data)\n                \n                loss = criterion(output, target)\n                loss.backward()\n                optimizer.step()\n                tepoch.set_postfix(loss=loss.item())\n                running_loss += loss.item()\n\n                # Get accuracy \n                _, predicted = torch.max(output, 1)\n                total += target.size(0)\n                correct += (predicted == target).sum().item()\n        \n        train_loss.append(running_loss/len(train_loader))\n        train_acc.append(correct/total)\n                \n        # Evaluate on validation data\n        model.eval()\n        running_loss = 0.\n        correct, total = 0, 0 \n        with tqdm(validation_loader, unit='batch') as tepoch:\n            tepoch.set_description('Validation: ')\n            for data, target in tepoch:\n                data, target = data.to(device), target.to(device)\n                optimizer.zero_grad()\n                output = model(data)\n                \n                loss = criterion(output, target)\n                tepoch.set_postfix(loss=loss.item())\n                running_loss += loss.item()\n\n                # Get accuracy \n                _, predicted = torch.max(output, 1)\n                total += target.size(0)\n                correct += (predicted == target).sum().item()\n        \n        validation_loss.append(running_loss/len(validation_loader))\n        validation_acc.append(correct/total)\n    \n    return train_loss, train_acc, validation_loss, validation_acc ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compiling and Training\n\nnet = Baseline_CNN().to(device)\ntrain_loss, train_acc, validation_loss, validation_acc = train(net, device, train_loader, val_loader, 20)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test model","metadata":{}},{"cell_type":"code","source":"torch.save(net.state_dict(), \"Transfer_Learning_Resnet.pt\")\nplot_loss_accuracy(train_loss, train_acc, validation_loss, validation_acc)\nprint(\"The Final Validation Accuracy was: \", \"{:.2f}\".format(100*validation_acc[-1]), \"%\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def print_saliency(i):\n    image_test = images[i]\n    # from: https://towardsdatascience.com/saliency-map-using-pytorch-68270fe45e80\n    image_test = image_test.reshape(1, 3, 224, 224)\n    image_test = image_test.to(device)\n    image_test.requires_grad_()\n\n    # Retrieve output from the image\n    output = net(image_test)\n\n    # Catch the output\n    output_idx = output.argmax()\n    output_max = output[0, output_idx]\n\n    # Do backpropagation to get the derivative of the output based on the image\n    output_max.backward()\n    \n    # Retireve the saliency map and also pick the maximum value from channels on each pixel.\n    # In this case, we look at dim=1. Recall the shape (batch_size, channel, width, height)\n    saliency, _ = torch.max(image_test.grad.data.abs(), dim=1) \n    saliency = saliency.reshape(224, 224)\n\n    # Reshape the image\n    image_test = image_test.reshape(-1, 224, 224)\n\n    # Visualize the image and the saliency map\n    fig, ax = plt.subplots(1, 2)\n    ax[0].imshow(image_test.cpu().detach().numpy().transpose(1, 2, 0))\n    ax[0].axis('off')\n    ax[1].imshow(saliency.cpu(), cmap='hot')\n    ax[1].axis('off')\n    plt.tight_layout()\n    fig.suptitle('The Image and Its Saliency Map')\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print_saliency(111)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"while True: continue","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print_saliency(105)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print_saliency(15)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}