{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Imports here\nfrom __future__ import print_function, division\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torch.utils import data\nimport torch\nfrom torch import nn\nfrom torch import optim\nimport torchvision\nimport torch.nn.functional as F\nfrom torchvision import datasets, transforms, models\nimport torchvision.models as models\nfrom torch.utils.data.sampler import SubsetRandomSampler\nfrom torch.utils.data import Dataset, DataLoader\nfrom skimage import io, transform\nimport torch.utils.data as data_utils\nfrom PIL import Image, ImageFile\nimport json\nfrom torch.optim import lr_scheduler\nimport time\nimport os\nimport argparse\nimport copy\nimport pandas as pd\nImageFile.LOAD_TRUNCATED_IMAGES = True\nimport cv2\n# Import useful sklearn functions\nimport sklearn\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\n\nimport time\nfrom tqdm import tqdm_notebook\n\nimport os\nprint(os.listdir(\"../input\"))\nbase_dir = \"../input/aptos2019-blindness-detection/\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-10T09:23:03.108685Z","iopub.execute_input":"2023-09-10T09:23:03.109554Z","iopub.status.idle":"2023-09-10T09:23:09.629965Z","shell.execute_reply.started":"2023-09-10T09:23:03.109517Z","shell.execute_reply":"2023-09-10T09:23:09.629044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(os.listdir(\"../input/aptos2019-blindness-detection/\"))","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:25:03.126266Z","iopub.execute_input":"2023-09-10T09:25:03.126960Z","iopub.status.idle":"2023-09-10T09:25:03.134792Z","shell.execute_reply.started":"2023-09-10T09:25:03.126924Z","shell.execute_reply":"2023-09-10T09:25:03.133402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv = pd.read_csv('../input/aptos2019-blindness-detection/train.csv')\ntest_csv = pd.read_csv('../input/aptos2019-blindness-detection/test.csv')","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:25:03.446201Z","iopub.execute_input":"2023-09-10T09:25:03.447103Z","iopub.status.idle":"2023-09-10T09:25:03.482351Z","shell.execute_reply.started":"2023-09-10T09:25:03.447063Z","shell.execute_reply":"2023-09-10T09:25:03.481178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Train Size = {}'.format(len(train_csv)))\nprint('Public Test Size = {}'.format(len(test_csv)))","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:25:03.735244Z","iopub.execute_input":"2023-09-10T09:25:03.736415Z","iopub.status.idle":"2023-09-10T09:25:03.742657Z","shell.execute_reply.started":"2023-09-10T09:25:03.736369Z","shell.execute_reply":"2023-09-10T09:25:03.741216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv.head()","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:25:04.289205Z","iopub.execute_input":"2023-09-10T09:25:04.290295Z","iopub.status.idle":"2023-09-10T09:25:04.312623Z","shell.execute_reply.started":"2023-09-10T09:25:04.290259Z","shell.execute_reply":"2023-09-10T09:25:04.311079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:25:04.314624Z","iopub.execute_input":"2023-09-10T09:25:04.315008Z","iopub.status.idle":"2023-09-10T09:25:04.685600Z","shell.execute_reply.started":"2023-09-10T09:25:04.314977Z","shell.execute_reply":"2023-09-10T09:25:04.684207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\n\ncounts = train_csv['diagnosis'].value_counts()\nclass_list = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferate']\nfor i, x in enumerate(class_list):\n    counts[x] = counts.pop(i)\n\nplt.figure(figsize=(10, 5))\nsns.barplot(x=counts.index, y=counts.values, alpha=0.8, palette='bright')\nplt.title('Distribution of Output Classes')\nplt.ylabel('Number of Occurrences', fontsize=12)\nplt.xlabel('Target Classes', fontsize=12)\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:25:04.687810Z","iopub.execute_input":"2023-09-10T09:25:04.688171Z","iopub.status.idle":"2023-09-10T09:25:05.059792Z","shell.execute_reply.started":"2023-09-10T09:25:04.688141Z","shell.execute_reply":"2023-09-10T09:25:05.058496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Visualizing Training DataVisualizing Training Data","metadata":{}},{"cell_type":"code","source":"fig = plt.figure(figsize=(30, 6))\n# display 20 images\ntrain_imgs = os.listdir(base_dir+\"/train_images\")\nfor idx, img in enumerate(np.random.choice(train_imgs, 16)):\n    ax = fig.add_subplot(2, 16//2, idx+1, xticks=[], yticks=[])\n    im = Image.open(base_dir+\"/train_images/\" + img)\n    plt.imshow(im)\n    lab = train_csv.loc[train_csv['id_code'] == img.split('.')[0], 'diagnosis'].values[0]\n    ax.set_title('Severity: %s'%lab)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:25:06.999910Z","iopub.execute_input":"2023-09-10T09:25:07.000394Z","iopub.status.idle":"2023-09-10T09:25:24.343055Z","shell.execute_reply.started":"2023-09-10T09:25:07.000358Z","shell.execute_reply":"2023-09-10T09:25:24.341804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Visualizing Test SetVisualizing Test Set","metadata":{}},{"cell_type":"code","source":"fig = plt.figure(figsize=(30, 6))\n# display 20 images\ntest_imgs = os.listdir(base_dir+\"/test_images\")\nfor idx, img in enumerate(np.random.choice(test_imgs, 16)):\n    ax = fig.add_subplot(2, 16//2, idx+1, xticks=[], yticks=[])\n    im = Image.open(base_dir+\"/test_images/\" + img)\n    plt.imshow(im)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:25:45.442645Z","iopub.execute_input":"2023-09-10T09:25:45.443093Z","iopub.status.idle":"2023-09-10T09:25:53.395628Z","shell.execute_reply.started":"2023-09-10T09:25:45.443056Z","shell.execute_reply":"2023-09-10T09:25:53.394540Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Our own custom class for datasets\nclass CreateDataset(Dataset):\n    def __init__(self, df_data, data_dir = '../input/', transform=None):\n        super().__init__()\n        self.df = df_data.values\n        self.data_dir = data_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_name,label = self.df[index]\n        img_path = os.path.join(self.data_dir, img_name+'.png')\n        image = cv2.imread(img_path)\n        if self.transform is not None:\n            image = self.transform(image)\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:26:51.928566Z","iopub.execute_input":"2023-09-10T09:26:51.929033Z","iopub.status.idle":"2023-09-10T09:26:51.939366Z","shell.execute_reply.started":"2023-09-10T09:26:51.928998Z","shell.execute_reply":"2023-09-10T09:26:51.937661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.4),\n    #transforms.ColorJitter(brightness=2, contrast=2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225))\n])","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:20.785693Z","iopub.execute_input":"2023-09-10T09:27:20.786100Z","iopub.status.idle":"2023-09-10T09:27:20.793324Z","shell.execute_reply.started":"2023-09-10T09:27:20.786070Z","shell.execute_reply":"2023-09-10T09:27:20.791497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_transforms = transforms.Compose([transforms.Resize(256),\n                                      transforms.CenterCrop(224),\n                                      transforms.ToTensor(),\n                                      transforms.Normalize([0.485, 0.456, 0.406],[0.229, 0.224, 0.225])])","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:21.669565Z","iopub.execute_input":"2023-09-10T09:27:21.669996Z","iopub.status.idle":"2023-09-10T09:27:21.676877Z","shell.execute_reply.started":"2023-09-10T09:27:21.669964Z","shell.execute_reply":"2023-09-10T09:27:21.675268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = \"../input/aptos2019-blindness-detection/train_images/\"\ntest_path = \"../input/aptos2019-blindness-detection/test_images/\"","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:22.335833Z","iopub.execute_input":"2023-09-10T09:27:22.336224Z","iopub.status.idle":"2023-09-10T09:27:22.342032Z","shell.execute_reply.started":"2023-09-10T09:27:22.336194Z","shell.execute_reply":"2023-09-10T09:27:22.340700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = CreateDataset(df_data=train_csv, data_dir=train_path, transform=train_transforms)\ntest_data = CreateDataset(df_data=test_csv, data_dir=test_path, transform=test_transforms)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:23.014925Z","iopub.execute_input":"2023-09-10T09:27:23.015347Z","iopub.status.idle":"2023-09-10T09:27:23.021516Z","shell.execute_reply.started":"2023-09-10T09:27:23.015315Z","shell.execute_reply":"2023-09-10T09:27:23.020569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_size = 0.2\nnum_train = len(train_data)\nindices = list(range(num_train))\nnp.random.shuffle(indices)\nsplit = int(np.floor(valid_size * num_train))\ntrain_idx, valid_idx = indices[split:], indices[:split]\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:23.996049Z","iopub.execute_input":"2023-09-10T09:27:23.997073Z","iopub.status.idle":"2023-09-10T09:27:24.004925Z","shell.execute_reply.started":"2023-09-10T09:27:23.997018Z","shell.execute_reply":"2023-09-10T09:27:24.003580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_sampler = SubsetRandomSampler(train_idx)\nvalid_sampler = SubsetRandomSampler(valid_idx)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:24.628045Z","iopub.execute_input":"2023-09-10T09:27:24.628424Z","iopub.status.idle":"2023-09-10T09:27:24.633744Z","shell.execute_reply.started":"2023-09-10T09:27:24.628394Z","shell.execute_reply":"2023-09-10T09:27:24.632575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainloader = torch.utils.data.DataLoader(train_data, batch_size=64,sampler=train_sampler)\nvalidloader = torch.utils.data.DataLoader(train_data, batch_size=64, sampler=valid_sampler)\ntestloader = torch.utils.data.DataLoader(test_data, batch_size=64)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:25.253494Z","iopub.execute_input":"2023-09-10T09:27:25.253877Z","iopub.status.idle":"2023-09-10T09:27:25.260526Z","shell.execute_reply.started":"2023-09-10T09:27:25.253848Z","shell.execute_reply":"2023-09-10T09:27:25.259167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"training examples contain : {len(train_data)}\")\nprint(f\"testing examples contain : {len(test_data)}\")\n\nprint(len(trainloader))\nprint(len(validloader))\nprint(len(testloader))","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:26.609285Z","iopub.execute_input":"2023-09-10T09:27:26.609769Z","iopub.status.idle":"2023-09-10T09:27:26.619245Z","shell.execute_reply.started":"2023-09-10T09:27:26.609734Z","shell.execute_reply":"2023-09-10T09:27:26.617775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plotting the images of loaded batch with given fig size and frame data    \nimport torchvision\nimport matplotlib.pyplot as plt\nimport numpy as np\ngrid = torchvision.utils.make_grid(images, nrow = 20, padding = 2)\nplt.figure(figsize = (20, 20))  \nplt.imshow(np.transpose(grid, (1, 2, 0)))   \nprint('labels:', labels)  ","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:27.780728Z","iopub.execute_input":"2023-09-10T09:27:27.781136Z","iopub.status.idle":"2023-09-10T09:27:29.514012Z","shell.execute_reply.started":"2023-09-10T09:27:27.781106Z","shell.execute_reply":"2023-09-10T09:27:29.512892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# LOAD ONE BATCH OF TESTING SET TO CHECK THE IMAGES AND THEIR LABELS\nimages, labels = next(iter(trainloader))\n\n# Checking shape of image\nprint(f\"Image shape : {images.shape}\")\nprint(f\"Label shape : {labels.shape}\")\n\n# denormalizing images\ndef imshow(inp, title=None):\n    \"\"\"Imshow for Tensor.\"\"\"\n    inp = inp.numpy().transpose((1, 2, 0))\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    inp = std * inp + mean\n    inp = np.clip(inp, 0, 1)\n    plt.imshow(inp)\n    if title is not None:\n        plt.title(title)\n    plt.pause(0.001)","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:40.108762Z","iopub.execute_input":"2023-09-10T09:27:40.109174Z","iopub.status.idle":"2023-09-10T09:27:50.954748Z","shell.execute_reply.started":"2023-09-10T09:27:40.109143Z","shell.execute_reply":"2023-09-10T09:27:50.953510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plotting the images of loaded batch with given fig size and frame data    \nimport torchvision\nimport matplotlib.pyplot as plt\nimport numpy as np\ngrid = torchvision.utils.make_grid(images, nrow = 20, padding = 2)\nplt.figure(figsize = (20, 20))  \nplt.imshow(np.transpose(grid, (1, 2, 0)))   \nprint('labels:', labels)  ","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:50.957221Z","iopub.execute_input":"2023-09-10T09:27:50.957732Z","iopub.status.idle":"2023-09-10T09:27:52.695810Z","shell.execute_reply.started":"2023-09-10T09:27:50.957698Z","shell.execute_reply":"2023-09-10T09:27:52.694488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_names = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative DR']\n\nimages, labels = next(iter(trainloader))\nout = torchvision.utils.make_grid(images)\nimshow(out, title=[class_names[x] for x in labels])","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:27:58.535927Z","iopub.execute_input":"2023-09-10T09:27:58.536358Z","iopub.status.idle":"2023-09-10T09:28:10.447679Z","shell.execute_reply.started":"2023-09-10T09:27:58.536327Z","shell.execute_reply":"2023-09-10T09:28:10.446409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_on_gpu = torch.cuda.is_available()\n\nif not train_on_gpu:\n    print('CUDA is not available.  Training on CPU ...')\nelse:\n    print('CUDA is available!  Training on GPU ...')","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:28:19.641649Z","iopub.execute_input":"2023-09-10T09:28:19.642051Z","iopub.status.idle":"2023-09-10T09:28:19.649808Z","shell.execute_reply.started":"2023-09-10T09:28:19.642011Z","shell.execute_reply":"2023-09-10T09:28:19.648572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install timm\n\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:28:20.844666Z","iopub.execute_input":"2023-09-10T09:28:20.845186Z","iopub.status.idle":"2023-09-10T09:28:36.469718Z","shell.execute_reply.started":"2023-09-10T09:28:20.845153Z","shell.execute_reply":"2023-09-10T09:28:36.468383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install vit-pytorch\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:28:36.472761Z","iopub.execute_input":"2023-09-10T09:28:36.474019Z","iopub.status.idle":"2023-09-10T09:28:51.720270Z","shell.execute_reply.started":"2023-09-10T09:28:36.473959Z","shell.execute_reply":"2023-09-10T09:28:51.719110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport timm\nfrom torch.optim import lr_scheduler\nfrom vit_pytorch.cct import CCT\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Create the first model (CCT)\ncct = CCT(\n    img_size=(224, 224),\n    embedding_dim=384,\n    n_conv_layers=2,\n    kernel_size=7,\n    stride=2,\n    padding=3,\n    pooling_kernel_size=3,\n    pooling_stride=2,\n    pooling_padding=1,\n    num_layers=14,\n    num_heads=6,\n    mlp_ratio=3.,\n    num_classes=5,  # Adjust the number of classes for your task\n    positional_embedding='learnable'  # Choose the positional embedding type\n)\n\n# Load the pretrained CCT model weights if available\n# cct.load_state_dict(torch.load('pretrained_cct_model.pth'))\n\n# Create the second model (Vision Transformer)\nvit = timm.create_model('vit_base_patch16_224', pretrained=True)\nnum_ftrs = vit.head.in_features  # Get the number of input features of the model head\nout_ftrs = 5  # Adjust the number of classes for your task\nvit.head = nn.Sequential(\n    nn.Linear(num_ftrs, 512),\n    nn.ReLU(),\n    nn.Linear(512, out_ftrs),\n    nn.LogSoftmax(dim=1)\n)\n\n# Load the pretrained Vision Transformer model weights if available\n# vit.load_state_dict(torch.load('pretrained_vit_model.pth'))\n\n# Combine the two models into an ensemble\nclass EnsembleModel(nn.Module):\n    def __init__(self, model1, model2):\n        super(EnsembleModel, self).__init__()\n        self.model1 = model1\n        self.model2 = model2\n\n    def forward(self, x):\n        output1 = self.model1(x)\n        output2 = self.model2(x)\n        output = (output1 + output2) / 2.0  # Combine predictions using averaging\n        return output\n\n# Create the ensemble model\nensemble = EnsembleModel(cct, vit)\n\n# Define loss function, optimizer, and scheduler\ncriterion = nn.NLLLoss()\noptimizer = optim.Adam(ensemble.parameters(), lr=0.00001)\nscheduler = lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n\nensemble.to(device)\n\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:28:51.723091Z","iopub.execute_input":"2023-09-10T09:28:51.723502Z","iopub.status.idle":"2023-09-10T09:29:03.029166Z","shell.execute_reply.started":"2023-09-10T09:28:51.723454Z","shell.execute_reply":"2023-09-10T09:29:03.027743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.Adam(filter(lambda p: p.requires_grad, ensemble.parameters()), lr=0.000001)\nscheduler = lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:29:12.436287Z","iopub.execute_input":"2023-09-10T09:29:12.436713Z","iopub.status.idle":"2023-09-10T09:29:12.446299Z","shell.execute_reply.started":"2023-09-10T09:29:12.436683Z","shell.execute_reply":"2023-09-10T09:29:12.444952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pytorch_total_params = sum(p.numel() for p in ensemble.parameters() if p.requires_grad)\nprint(\"Number of trainable parameters: \\n{}\".format(pytorch_total_params))\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:29:13.330263Z","iopub.execute_input":"2023-09-10T09:29:13.330730Z","iopub.status.idle":"2023-09-10T09:29:13.342355Z","shell.execute_reply.started":"2023-09-10T09:29:13.330694Z","shell.execute_reply":"2023-09-10T09:29:13.340113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_losses, valid_losses, acc = train_and_test(5, ensemble)\n","metadata":{"execution":{"iopub.status.busy":"2023-09-10T09:29:17.001538Z","iopub.execute_input":"2023-09-10T09:29:17.003356Z","iopub.status.idle":"2023-09-10T16:16:22.813954Z","shell.execute_reply.started":"2023-09-10T09:29:17.003299Z","shell.execute_reply":"2023-09-10T16:16:22.806331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}],"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"}}