{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "cc7a778d",
   "metadata": {},
   "source": [
    "## Imports"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a89d731b",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "from torchvision.datasets import ImageFolder\n",
    "from collections import Counter\n",
    "from torchvision import transforms\n",
    "import matplotlib.pyplot as plt\n",
    "from torchvision.utils import make_grid\n",
    "import numpy as np\n",
    "from functools import partial"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "eb6b3400",
   "metadata": {},
   "source": [
    "## Preprocessing"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cbd523d9",
   "metadata": {},
   "outputs": [],
   "source": [
    "mean = [0.1918]\n",
    "std = [0.2148]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8a364877",
   "metadata": {},
   "outputs": [],
   "source": [
    "transform = transforms.Compose([\n",
    "    transforms.Grayscale(num_output_channels=1),\n",
    "    transforms.ToTensor(),\n",
    "    transforms.Normalize(mean, std)\n",
    "])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a0112fdf",
   "metadata": {},
   "outputs": [],
   "source": [
    "dataset = ImageFolder('../dataset', transform)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0fb1b47e",
   "metadata": {},
   "outputs": [],
   "source": [
    "count_dict = dict(Counter(dataset.targets))\n",
    "count = count_dict.values()\n",
    "total = sum(count)\n",
    "weight = [total/c for c in count]\n",
    "weight = torch.FloatTensor(weight)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f3661db5",
   "metadata": {},
   "outputs": [],
   "source": [
    "dataset_size = len(dataset)\n",
    "train_size = int(dataset_size * 0.8)\n",
    "val_size = int(dataset_size * 0.1)\n",
    "test_size = dataset_size - (train_size + val_size)\n",
    "\n",
    "train_dataset, val_dataset, test_dataset = torch.utils.data.random_split(dataset, [train_size, val_size, test_size], generator=torch.Generator().manual_seed(42))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "082727cc",
   "metadata": {},
   "outputs": [],
   "source": [
    "def create_data_loaders(train_dataset, val_dataset, test_dataset, batch_size):\n",
    "    train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\n",
    "    val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=batch_size, shuffle=True)\n",
    "    test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size, shuffle=True)\n",
    "    \n",
    "    return train_loader, val_loader, test_loader"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "dd2cc4c7",
   "metadata": {},
   "outputs": [],
   "source": [
    "def show_batch(dl):\n",
    "    for images, labels in dl:\n",
    "        fig,ax = plt.subplots(figsize = (8,8))\n",
    "        ax.set_xticks([])\n",
    "        ax.set_yticks([])\n",
    "        ax.imshow(make_grid(images.mul_(torch.as_tensor(mean)).add_(torch.as_tensor(std)), nrow=4, pad_value=1).permute(1,2,0))\n",
    "        break"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f711d42f",
   "metadata": {},
   "outputs": [],
   "source": [
    "def calc_mean_std(data_loader):\n",
    "    channels_sum, channels_sqrd_sum, num_batches = 0, 0, 0\n",
    "\n",
    "    for data, _ in data_loader:\n",
    "        channels_sum += torch.mean(data, dim=[0, 2, 3])\n",
    "        channels_sqrd_sum += torch.mean(data ** 2, dim=[0, 2, 3])\n",
    "        num_batches += 1\n",
    "\n",
    "    mean = channels_sum / num_batches\n",
    "    std = (channels_sqrd_sum / num_batches - mean ** 2) ** 0.5\n",
    "\n",
    "    return mean, std"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "4b531f45",
   "metadata": {},
   "source": [
    "# Models"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8e78ba08",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch.nn as nn\n",
    "import torch.nn.functional as F\n",
    "\n",
    "class ImageClassificationBase(nn.Module):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "        self._initialize_weights()\n",
    "    \n",
    "    def training_step(self, batch, criterion):\n",
    "        images, labels = batch \n",
    "        out = self(images) \n",
    "        loss = criterion(out, labels) \n",
    "        acc = accuracy(out, labels)\n",
    "        return loss, acc\n",
    "    \n",
    "    def validation_step(self, batch, criterion):\n",
    "        images, labels = batch \n",
    "        out = self(images)                  \n",
    "        loss = criterion(out, labels)  \n",
    "        acc = accuracy(out, labels)          \n",
    "        return {'val_loss': loss.detach(), 'val_acc': acc}\n",
    "        \n",
    "    def validation_epoch_end(self, outputs):\n",
    "        batch_losses = [x['val_loss'] for x in outputs]\n",
    "        epoch_loss = torch.stack(batch_losses).mean()   \n",
    "        batch_accs = [x['val_acc'] for x in outputs]\n",
    "        epoch_acc = torch.stack(batch_accs).mean()     \n",
    "        return {'val_loss': epoch_loss.item(), 'val_acc': epoch_acc.item()}\n",
    "    \n",
    "    def epoch_end(self, epoch, result):\n",
    "        print(\"Epoch [{}], train_loss: {:.4f}, val_loss: {:.4f}, train_acc: {:.4f}, val_acc: {:.4f}\".format(\n",
    "            epoch, result['train_loss'], result['val_loss'], result['train_acc'], result['val_acc']))\n",
    "        \n",
    "    def _initialize_weights(m):\n",
    "        if isinstance(m, nn.Conv2d) or isinstance(m, torch.nn.Linear):\n",
    "            nn.init.kaiming_normal_(w, nonlinearity='relu')\n",
    "            if m.bias is not None:\n",
    "                torch.nn.init.zeros_(m.bias)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2517a155",
   "metadata": {},
   "outputs": [],
   "source": [
    "class ConvBlock(nn.Module):\n",
    "    def __init__(self, in_channels, out_channels, **kwargs):\n",
    "        super().__init__()\n",
    "        self.block = nn.Sequential(\n",
    "            nn.Conv2d(in_channels, out_channels, **kwargs),\n",
    "            nn.BatchNorm2d(out_channels),\n",
    "            nn.ReLU()\n",
    "        )\n",
    "\n",
    "    def forward(self, x):\n",
    "        return self.block(x)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "04ba7ef9",
   "metadata": {},
   "source": [
    "## AlexNet"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d86cd2b5",
   "metadata": {},
   "outputs": [],
   "source": [
    "class AlexNet(ImageClassificationBase):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "        self.network = nn.Sequential(\n",
    "            nn.Conv2d(1, 16, kernel_size=11, stride=4, padding=2),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(kernel_size=3, stride=2),\n",
    "            nn.Conv2d(16, 32, kernel_size=5, stride=1, padding=2),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(kernel_size=3, stride=2),\n",
    "            nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1),\n",
    "            nn.ReLU(),\n",
    "            nn.AdaptiveAvgPool2d(output_size=6),\n",
    "            nn.Dropout(0.5),\n",
    "            nn.Flatten(),\n",
    "            nn.Linear(2304, 256, bias=True),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(0.5),\n",
    "            nn.Linear(256, 4)\n",
    "        )\n",
    "        \n",
    "    def forward(self, x):\n",
    "        return self.network(x)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fdc79869",
   "metadata": {},
   "outputs": [],
   "source": [
    "class AlexNet2(ImageClassificationBase):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "        self.network = nn.Sequential(\n",
    "            nn.Conv2d(1, 16, kernel_size=11, stride=4, padding=2),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(kernel_size=3, stride=2),\n",
    "            nn.Conv2d(16, 32, kernel_size=5, stride=1, padding=2),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(kernel_size=3, stride=2),\n",
    "            nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=2),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(kernel_size=3, stride=2),\n",
    "            nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1),\n",
    "            nn.ReLU(),\n",
    "            nn.AdaptiveAvgPool2d(output_size=4),\n",
    "            nn.Dropout(0.5),\n",
    "            nn.Flatten(),\n",
    "            nn.Linear(1024, 256, bias=True),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(0.5),\n",
    "            nn.Linear(256, 4)\n",
    "        )\n",
    "        \n",
    "    def forward(self, x):\n",
    "        return self.network(x)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "e812706a",
   "metadata": {},
   "source": [
    "## VGG"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "961fb74a",
   "metadata": {},
   "outputs": [],
   "source": [
    "class VGG1(ImageClassificationBase):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "        self.network = nn.Sequential(\n",
    "            \n",
    "            nn.Conv2d(1, 16, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.Conv2d(16, 32, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(2,2),\n",
    "        \n",
    "            nn.Conv2d(32, 64, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.Conv2d(64 ,64, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.AdaptiveAvgPool2d(output_size=7),\n",
    "            \n",
    "            nn.Flatten(),\n",
    "            nn.Linear(3136, 128),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(0.85),\n",
    "            nn.Linear(128, 64),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(0.85),\n",
    "            nn.Linear(64, 4)\n",
    "        )\n",
    "    \n",
    "    def forward(self, x):\n",
    "        return self.network(x)\n",
    "    \n",
    "    def _initialize_weights(m):\n",
    "        if isinstance(m, nn.Conv2d) or isinstance(m, torch.nn.Linear):\n",
    "            torch.nn.init.xavier_uniform_(m.weight)\n",
    "            if m.bias is not None:\n",
    "                torch.nn.init.zeros_(m.bias)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0ddd2181",
   "metadata": {},
   "outputs": [],
   "source": [
    "class VGG2(ImageClassificationBase):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "        self.network = nn.Sequential(\n",
    "            \n",
    "            nn.Conv2d(1, 16, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.Conv2d(16, 32, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(2,2),\n",
    "        \n",
    "            nn.Conv2d(32, 48, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.Conv2d(48 ,64, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.AdaptiveAvgPool2d(output_size=4),\n",
    "            \n",
    "            nn.Flatten(),\n",
    "            nn.Linear(1024, 128),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(0.85),\n",
    "            nn.Linear(128, 4)\n",
    "        )\n",
    "    \n",
    "    def forward(self, x):\n",
    "        return self.network(x)\n",
    "    \n",
    "\n",
    "    def _initialize_weights(m):\n",
    "        if isinstance(m, nn.Conv2d) or isinstance(m, torch.nn.Linear):\n",
    "            torch.nn.init.xavier_uniform_(m.weight)\n",
    "            if m.bias is not None:\n",
    "                torch.nn.init.zeros_(m.bias)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b2862394",
   "metadata": {},
   "outputs": [],
   "source": [
    "class VGG3(ImageClassificationBase):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "        self.network = nn.Sequential(\n",
    "            \n",
    "            nn.Conv2d(1, 16, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.Conv2d(16, 32, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(2,2),\n",
    "            \n",
    "            nn.Conv2d(32, 48, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.Conv2d(48, 64, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.MaxPool2d(2,2),\n",
    "        \n",
    "            nn.Conv2d(64, 96, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.Conv2d(96 ,96, kernel_size=3),\n",
    "            nn.ReLU(),\n",
    "            nn.AdaptiveAvgPool2d(output_size=3),\n",
    "            \n",
    "            nn.Flatten(),\n",
    "            nn.Linear(864, 128),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(0.85),\n",
    "            nn.Linear(128, 64),\n",
    "            nn.ReLU(),\n",
    "            nn.Dropout(0.85),\n",
    "            nn.Linear(64, 4)\n",
    "        )\n",
    "    \n",
    "    def forward(self, x):\n",
    "        return self.network(x)\n",
    "    \n",
    "    def _initialize_weights(m):\n",
    "        if isinstance(m, nn.Conv2d) or isinstance(m, torch.nn.Linear):\n",
    "            torch.nn.init.xavier_uniform_(m.weight)\n",
    "            if m.bias is not None:\n",
    "                torch.nn.init.zeros_(m.bias)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "fe1c228c",
   "metadata": {},
   "source": [
    "## Inception Net"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "37be7298",
   "metadata": {},
   "outputs": [],
   "source": [
    "class InceptionModuleV1(nn.Module):\n",
    "\n",
    "    def __init__(self, in_channels, out_channels: dict, reduction_channels: dict):\n",
    "\n",
    "        super().__init__()\n",
    "        \n",
    "        self.conv_1x1 = ConvBlock(in_channels, out_channels['1x1'], kernel_size=1)\n",
    "        \n",
    "        self.conv_3x3 = nn.Sequential(\n",
    "            ConvBlock(in_channels, reduction_channels['3x3'], kernel_size=1),\n",
    "            ConvBlock(reduction_channels['3x3'], out_channels['3x3'], kernel_size=3, padding=1)\n",
    "        )\n",
    "        \n",
    "        self.conv_5x5 = nn.Sequential(\n",
    "            ConvBlock(in_channels, reduction_channels['5x5'], kernel_size=1),\n",
    "            ConvBlock(reduction_channels['5x5'], out_channels['5x5'], kernel_size=5, padding=2)\n",
    "        )\n",
    "\n",
    "        self.max_pool = nn.Sequential(\n",
    "            nn.MaxPool2d(kernel_size=3, padding=1, stride=1),\n",
    "            ConvBlock(in_channels, out_channels['max'], kernel_size=1),\n",
    "            nn.ReLU()\n",
    "        )\n",
    "\n",
    "    def forward(self, x):\n",
    "        return torch.cat([self.conv_1x1(x), self.conv_3x3(x), self.conv_5x5(x), self.max_pool(x)], dim=1)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b56dd3ec",
   "metadata": {},
   "outputs": [],
   "source": [
    "class Inception1(ImageClassificationBase):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "\n",
    "        self.input_net = ConvBlock(1, 64, kernel_size=3, padding=1)\n",
    "        \n",
    "        self.inception_blocks = nn.Sequential(\n",
    "            InceptionModuleV1(64, out_channels={'1x1': 16, '3x3': 32, '5x5': 8, 'max': 8}, reduction_channels={'3x3': 32, '5x5': 16}),\n",
    "            InceptionModuleV1(64, out_channels={'1x1': 24, '3x3': 48, '5x5': 12, 'max': 12}, reduction_channels={'3x3': 32, '5x5': 16}),\n",
    "            nn.MaxPool2d(3, stride=2, padding=1), \n",
    "            InceptionModuleV1(96, out_channels={'1x1': 24, '3x3': 48, '5x5': 12, 'max': 12}, reduction_channels={'3x3': 32, '5x5': 16}),\n",
    "            InceptionModuleV1(96, out_channels={'1x1': 16, '3x3': 48, '5x5': 16, 'max': 16}, reduction_channels={'3x3': 32, '5x5': 16}),\n",
    "            InceptionModuleV1(96, out_channels={'1x1': 32, '3x3': 48, '5x5': 24, 'max': 24}, reduction_channels={'3x3': 32, '5x5': 16}),\n",
    "            nn.MaxPool2d(3, stride=2, padding=1),\n",
    "            InceptionModuleV1(128, out_channels={'1x1': 32, '3x3': 64, '5x5': 16, 'max': 16}, reduction_channels={'3x3': 48, '5x5': 16}),\n",
    "            InceptionModuleV1(128, out_channels={'1x1': 32, '3x3': 64, '5x5': 16, 'max': 16}, reduction_channels={'3x3': 48, '5x5': 16})\n",
    "        )\n",
    "\n",
    "        self.output_net = nn.Sequential(\n",
    "            nn.AdaptiveAvgPool2d((1, 1)),\n",
    "            nn.Flatten(),\n",
    "            nn.Linear(128, 64),\n",
    "            nn.Dropout(0.4),\n",
    "            nn.Linear(64, 4)\n",
    "        )\n",
    "\n",
    "\n",
    "    def forward(self, x):\n",
    "        x = self.input_net(x)\n",
    "        x = self.inception_blocks(x)\n",
    "        x = self.output_net(x)\n",
    "        return x"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2e061984",
   "metadata": {},
   "outputs": [],
   "source": [
    "class Inception2(ImageClassificationBase):\n",
    "    def __init__(self):\n",
    "        super().__init__()\n",
    "        \n",
    "        self.input_net = ConvBlock(1, 32, kernel_size=3, padding=1)\n",
    "        \n",
    "        self.inception_blocks = nn.Sequential(\n",
    "            InceptionModuleV1(32, out_channels={'1x1': 8, '3x3': 16, '5x5': 4, 'max': 4}, reduction_channels={'3x3': 12, '5x5': 4}),\n",
    "            InceptionModuleV1(32, out_channels={'1x1': 12, '3x3': 24, '5x5': 8, 'max': 8}, reduction_channels={'3x3': 16, '5x5': 8}),\n",
    "            nn.MaxPool2d(3, stride=2, padding=1), \n",
    "            InceptionModuleV1(52, out_channels={'1x1': 16, '3x3': 32, '5x5': 12, 'max': 12}, reduction_channels={'3x3': 16, '5x5': 8}),\n",
    "            InceptionModuleV1(72, out_channels={'1x1': 16, '3x3': 32, '5x5': 16, 'max': 16}, reduction_channels={'3x3': 24, '5x5': 12}),\n",
    "            InceptionModuleV1(80, out_channels={'1x1': 32, '3x3': 48, '5x5': 24, 'max': 24}, reduction_channels={'3x3': 32, '5x5': 16}),\n",
    "            nn.MaxPool2d(3, stride=2, padding=1),\n",
    "            InceptionModuleV1(128, out_channels={'1x1': 32, '3x3': 64, '5x5': 16, 'max': 16}, reduction_channels={'3x3': 48, '5x5': 16}),\n",
    "            InceptionModuleV1(128, out_channels={'1x1': 32, '3x3': 64, '5x5': 16, 'max': 16}, reduction_channels={'3x3': 48, '5x5': 16}),\n",
    "            nn.MaxPool2d(3, stride=2, padding=1),\n",
    "            InceptionModuleV1(128, out_channels={'1x1': 32, '3x3': 72, '5x5': 24, 'max': 24}, reduction_channels={'3x3': 48, '5x5': 16}),\n",
    "            InceptionModuleV1(152, out_channels={'1x1': 32, '3x3': 72, '5x5': 24, 'max': 24}, reduction_channels={'3x3': 48, '5x5': 16}),\n",
    "        )\n",
    "\n",
    "        self.output_net = nn.Sequential(\n",
    "            nn.AdaptiveAvgPool2d((1, 1)),\n",
    "            nn.Flatten(),\n",
    "            nn.Linear(152, 128),\n",
    "            nn.Dropout(0.4),\n",
    "            nn.Linear(128, 4)\n",
    "        )\n",
    "\n",
    "\n",
    "    def forward(self, x):\n",
    "        x = self.input_net(x)\n",
    "        x = self.inception_blocks(x)\n",
    "        x = self.output_net(x)\n",
    "        return x"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f51d94c8",
   "metadata": {},
   "source": [
    "## ResNet"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "44a6ca2e",
   "metadata": {},
   "outputs": [],
   "source": [
    "class ResidualBlock(nn.Module):\n",
    "    def __init__(self, in_channels, intermediate_channels, identity_downsample=None, stride=1):\n",
    "        super().__init__()\n",
    "        self.expansion = 4\n",
    "        self.blocks = nn.Sequential(\n",
    "            ConvBlock(in_channels, intermediate_channels, kernel_size=1, stride=1, padding=0, bias=False),\n",
    "            ConvBlock(intermediate_channels, intermediate_channels, kernel_size=3, stride=stride, padding=1, bias=False),\n",
    "            nn.Conv2d(intermediate_channels, intermediate_channels * self.expansion, kernel_size=1, stride=1, padding=0, bias=False),\n",
    "            nn.BatchNorm2d(intermediate_channels * self.expansion)\n",
    "        )\n",
    "        \n",
    "        self.relu = nn.ReLU()\n",
    "\n",
    "        self.identity_downsample = identity_downsample\n",
    "        self.stride = stride\n",
    "\n",
    "    def forward(self, x):\n",
    "        identity = x.clone()\n",
    "\n",
    "        x = self.blocks(x)\n",
    "\n",
    "        if self.identity_downsample is not None:\n",
    "            identity = self.identity_downsample(identity)\n",
    "\n",
    "        x += identity\n",
    "        x = self.relu(x)\n",
    "        return x\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "76c50d2a",
   "metadata": {},
   "outputs": [],
   "source": [
    "class ResNet(ImageClassificationBase):\n",
    "    def __init__(self, block, layers, image_channels=1, num_classes=4, apply_dropout=False):\n",
    "        super().__init__()\n",
    "        self.in_channels = 64\n",
    "        self.apply_dropout = apply_dropout\n",
    "        \n",
    "        self.network = nn.Sequential(\n",
    "            ConvBlock(image_channels, 64, kernel_size=7, stride=2, padding=3, bias=False),\n",
    "            nn.MaxPool2d(kernel_size=3, stride=2, padding=1),\n",
    "            self._make_layer(block, layers[0], intermediate_channels=16, stride=1),\n",
    "            self._make_layer(block, layers[1], intermediate_channels=32, stride=2),\n",
    "            self._make_layer(block, layers[2], intermediate_channels=64, stride=2),\n",
    "            nn.AdaptiveAvgPool2d((1, 1))\n",
    "        )\n",
    "\n",
    "        self.fc1 = nn.Linear(64 * 4, 64)\n",
    "        self.fc2 = nn.Linear(64, num_classes)\n",
    "    \n",
    "        self.dropout = nn.Dropout()\n",
    "\n",
    "    def forward(self, x):\n",
    "        x = self.network(x)\n",
    "        x = x.reshape(x.shape[0], -1)\n",
    "        if self.apply_dropout:\n",
    "            x = self.dropout(x)\n",
    "        x = self.fc1(x)\n",
    "        if self.apply_dropout:\n",
    "            x = self.dropout(x)\n",
    "        x = self.fc2(x)\n",
    "\n",
    "        return x\n",
    "\n",
    "    def _make_layer(self, block, num_residual_blocks, intermediate_channels, stride):\n",
    "        identity_downsample = None\n",
    "        layers = []\n",
    "\n",
    "        if stride != 1 or self.in_channels != intermediate_channels * 4:\n",
    "            identity_downsample = nn.Sequential(\n",
    "                nn.Conv2d(self.in_channels, intermediate_channels * 4, kernel_size=1, stride=stride, bias=False),\n",
    "                nn.BatchNorm2d(intermediate_channels * 4),\n",
    "            )\n",
    "\n",
    "        layers.append(\n",
    "            block(self.in_channels, intermediate_channels, identity_downsample, stride)\n",
    "        )\n",
    "\n",
    "        self.in_channels = intermediate_channels * 4\n",
    "\n",
    "        for i in range(num_residual_blocks - 1):\n",
    "            layers.append(block(self.in_channels, intermediate_channels))\n",
    "\n",
    "        return nn.Sequential(*layers)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "beecb851",
   "metadata": {},
   "source": [
    "## Training"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d57cd5c8",
   "metadata": {},
   "outputs": [],
   "source": [
    "import copy\n",
    "\n",
    "def accuracy(outputs, labels):\n",
    "    _, preds = torch.max(outputs.data, dim=1)\n",
    "    return torch.tensor(torch.sum(preds == labels).item() / len(preds))\n",
    "  \n",
    "@torch.no_grad()\n",
    "def evaluate(model, criterion, val_loader):\n",
    "    model.eval()\n",
    "    outputs = [model.validation_step(batch, criterion) for batch in val_loader]\n",
    "    return model.validation_epoch_end(outputs)\n",
    "\n",
    "  \n",
    "def fit(epochs, model, optimizer, criterion, train_loader, val_loader):\n",
    "    \n",
    "    history = []\n",
    "\n",
    "    # early stopping params\n",
    "    n_epochs_stop = 20\n",
    "    epochs_no_improve = 0\n",
    "    min_val_loss = None\n",
    "    best_model = None\n",
    "    \n",
    "    for epoch in range(epochs):\n",
    "        print(f'start epoch {epoch}')\n",
    "        \n",
    "        model.train()\n",
    "        train_losses = []\n",
    "        train_accuracies = []\n",
    "        for batch in train_loader:\n",
    "            optimizer.zero_grad()\n",
    "            loss, acc = model.training_step(batch, criterion)\n",
    "            train_losses.append(loss)\n",
    "            train_accuracies.append(acc)\n",
    "            loss.backward()\n",
    "            optimizer.step()\n",
    "\n",
    "        result = evaluate(model, criterion, val_loader)\n",
    "        result['train_loss'] = torch.stack(train_losses).mean().item()\n",
    "        result['train_acc'] = torch.stack(train_accuracies).mean().item()\n",
    "        model.epoch_end(epoch, result)\n",
    "        history.append(result)\n",
    "        \n",
    "        # early stopping\n",
    "        if min_val_loss == None or result['val_loss'] < min_val_loss:\n",
    "            epochs_no_improve = 0\n",
    "            min_val_loss = result['val_loss']\n",
    "            best_model = copy.deepcopy(model)\n",
    "        else:\n",
    "            epochs_no_improve += 1\n",
    "            if epochs_no_improve == n_epochs_stop:\n",
    "                print('Early stopping' )\n",
    "                break\n",
    "        \n",
    "    \n",
    "    return history, model, best_model"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "77440d19",
   "metadata": {},
   "outputs": [],
   "source": [
    "def train(model_name, opt_name, criterion, batch_size, lr, momentum, weight_decay, train_loader, val_loader):    \n",
    "    \n",
    "    params = model_params[model_name]\n",
    "    model = models[model_name](**params)\n",
    "    \n",
    "    optim = opt_func[opt_name]\n",
    "    optim_params = {'lr': lr, 'weight_decay': weight_decay}\n",
    "        \n",
    "    if opt_name == 'RMSprop':\n",
    "        optim_params['momentum'] = momentum\n",
    "\n",
    "    optimizer = optim(model.parameters(), **optim_params)\n",
    "    \n",
    "    return fit(num_epochs, model, optimizer, criterion, train_loader, val_loader)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "380a1ca3",
   "metadata": {},
   "source": [
    "## Evaluation"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ca49c256",
   "metadata": {},
   "outputs": [],
   "source": [
    "import pickle\n",
    "import json\n",
    "\n",
    "def save_obj(obj, name):\n",
    "    obj = json.dumps(obj)\n",
    "    f = open(name + '.json', 'w')\n",
    "    f.write(obj)\n",
    "    f.close()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "dec9076f",
   "metadata": {},
   "outputs": [],
   "source": [
    "def plot_history(path, id, history, train_str, val_str, y_label, title):\n",
    "    train = [x[train_str] for x in history]\n",
    "    val = [x[val_str] for x in history]\n",
    "    plt.figure().clear()\n",
    "    plt.plot(train, color='#4c9ac7', linestyle='solid', marker='x')\n",
    "    plt.plot(val, color='#de9d35', linestyle='solid', marker='.')\n",
    "    plt.xlabel('epoch')\n",
    "    plt.ylabel(y_label)\n",
    "    plt.title(title);\n",
    "    plt.savefig(f'{path}/{id}_{y_label}.png')\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bb0edbf5",
   "metadata": {},
   "outputs": [],
   "source": [
    "from sklearn.metrics import confusion_matrix\n",
    "import seaborn as sn\n",
    "import pandas as pd\n",
    "\n",
    "def get_confusion_matrix(path, id, model, test_loader):\n",
    "\n",
    "    y_pred = []\n",
    "    y_true = []\n",
    "\n",
    "    for images, labels in test_loader:\n",
    "            output = model(images)\n",
    "\n",
    "            _, preds = torch.max(output.data, dim=1)\n",
    "            y_pred.extend(preds.cpu().numpy()) \n",
    "\n",
    "            labels = labels.data.cpu().numpy()\n",
    "            y_true.extend(labels) \n",
    "\n",
    "    classes = ('cnv', 'dme', 'drusen', 'normal')\n",
    "\n",
    "    cf_matrix = confusion_matrix(y_true, y_pred, normalize='true')\n",
    "\n",
    "    df_cm = pd.DataFrame(cf_matrix, index = [i for i in classes], columns = [i for i in classes])\n",
    "\n",
    "    plt.figure().clear()\n",
    "    plt.figure(figsize = (12,7))\n",
    "    sn.heatmap(df_cm, annot=True, cmap=\"Blues\")\n",
    "    plt.savefig(f'{path}/{id}_conf_matrix.png')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d54232ce",
   "metadata": {},
   "outputs": [],
   "source": [
    "from sklearn.metrics import precision_recall_fscore_support, roc_auc_score, accuracy_score\n",
    "\n",
    "def get_precision_recall_f1(path, id, model, test_loader):\n",
    "    y_pred = []\n",
    "    y_true = []\n",
    "    y_score = []\n",
    "    softmax = nn.Softmax(dim=1)\n",
    "    res = {}\n",
    "\n",
    "    for inputs, labels in test_loader:\n",
    "        output = model(inputs)\n",
    "\n",
    "        _, preds = torch.max(output.data, dim=1)\n",
    "        y_pred.extend(preds.cpu().numpy()) \n",
    "        y_score.extend(softmax(output.data).cpu().numpy())\n",
    "\n",
    "        labels = labels.data.cpu().numpy()\n",
    "        y_true.extend(labels)\n",
    "\n",
    "    res['metrics'] = precision_recall_fscore_support(y_true, y_pred, average='macro')\n",
    "    res['roc_auc_score_ovo'] = roc_auc_score(y_true, y_score, multi_class='ovo')\n",
    "    res['roc_auc_score_ovr'] = roc_auc_score(y_true, y_score, multi_class='ovr')\n",
    "    res['accuracy'] = accuracy_score(y_true, y_pred)\n",
    "    save_obj(res, f'{path}/{id}_metrics')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a4cdf047",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "\n",
    "def log_model(id, history, model, best_model, test_loader):\n",
    "    path = f'models/{id}'\n",
    "    os.mkdir(path)\n",
    "    \n",
    "    model.eval()\n",
    "    best_model.eval()\n",
    "    \n",
    "    torch.save(model.state_dict(), f'{path}/{id}.pth')\n",
    "    torch.save(best_model.state_dict(), f'{path}/{id}_best.pth')\n",
    "    save_obj(history, f'{path}/{id}_history')\n",
    "\n",
    "    get_confusion_matrix(path, id, model, test_loader)\n",
    "    get_confusion_matrix(path, f'{id}_best', best_model, test_loader)\n",
    "    \n",
    "    get_precision_recall_f1(path, id, model, test_loader)\n",
    "    get_precision_recall_f1(path, f'{id}_best', best_model, test_loader)\n",
    "    \n",
    "    plot_history(path, id, history, 'train_acc', 'val_acc', 'accuracy', 'Accuracy vs. No. of epochs')\n",
    "    plot_history(path, id, history, 'train_loss', 'val_loss', 'loss', 'Loss vs. No. of epochs')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "e94e411e",
   "metadata": {},
   "source": [
    "## Hyperparameters"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "eb77f45d",
   "metadata": {},
   "outputs": [],
   "source": [
    "batch_sizes = [16, 32]\n",
    "\n",
    "opt_func = {\n",
    "    'Adam': torch.optim.Adam, \n",
    "    'RMSprop': torch.optim.RMSprop\n",
    "}\n",
    "\n",
    "criterion = nn.CrossEntropyLoss(weight=weight)\n",
    "\n",
    "models = {\n",
    "    'AlexNet1': AlexNet,\n",
    "    'AlexNet2': AlexNet2,\n",
    "    'VGG1': VGG1,\n",
    "    'VGG2': VGG2,\n",
    "    'VGG3': VGG3,\n",
    "    'Inception1': Inception1,\n",
    "    'Inception2': Inception2,\n",
    "    'ResNet': ResNet\n",
    "}\n",
    "\n",
    "model_params = {\n",
    "    'AlexNet1': {},\n",
    "    'AlexNet2': {},\n",
    "    'VGG1': {},\n",
    "    'VGG2': {},\n",
    "    'VGG3': {},\n",
    "    'Inception1': {},\n",
    "    'Inception2': {},\n",
    "    'ResNet': {\n",
    "        'block': ResidualBlock,\n",
    "        'layers': [4,4,4],\n",
    "        'apply_dropout': True\n",
    "    }\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "05ff2bea",
   "metadata": {},
   "outputs": [],
   "source": [
    "model_name = 'ResNet'\n",
    "opt_name = 'RMSprop'\n",
    "batch_size = batch_sizes[1]\n",
    "num_epochs = 200\n",
    "lr = 0.00001\n",
    "momentum = 0.99\n",
    "weight_decay = 0.001\n",
    "\n",
    "id = f'{model_name}_{opt_name}_{batch_size}_lr{lr}' + (f'_m{momentum}' if opt_name == 'RMSprop' else '') + f'_wd{weight_decay}'\n",
    "\n",
    "train_loader, val_loader, test_loader = create_data_loaders(train_dataset, val_dataset, test_dataset, batch_size)\n",
    "history, model, best_model = train(model_name, opt_name, criterion, batch_size, lr, momentum, weight_decay, train_loader, val_loader)\n",
    "log_model(id, history, model, best_model, test_loader)"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.8.8"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
