{"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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport timm\nfrom torchvision import models\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom PIL import Image\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/aptos2019-blindness-detection\"\nIMG_DIR = os.path.join(DATA_DIR, \"train_images\")\nCSV_PATH = os.path.join(DATA_DIR, \"train.csv\")\n\ndf = pd.read_csv(CSV_PATH)\ndf[\"image_path\"] = df[\"id_code\"].apply(lambda x: os.path.join(IMG_DIR, x + \".png\"))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df, temp_df = train_test_split(\n    df, test_size=0.30, stratify=df[\"diagnosis\"], random_state=42\n)\n\nval_df, test_df = train_test_split(\n    temp_df, test_size=0.50, stratify=temp_df[\"diagnosis\"], random_state=42\n)\n\nprint(len(train_df), len(val_df), len(test_df))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_tfms = A.Compose([\n    A.RandomResizedCrop(\n        size=(224, 224),\n        scale=(0.8, 1.0),\n        ratio=(0.75, 1.33),\n        p=1.0\n    ),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.Rotate(limit=30, p=0.5),\n    A.RandomBrightnessContrast(p=0.4),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406),\n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2()\n])\n\nval_tfms = A.Compose([\n    A.Resize(224, 224),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406),\n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2()\n])\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class APTOSDataset(Dataset):\n    def __init__(self, df, transform):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img = Image.open(self.df.loc[idx, \"image_path\"]).convert(\"RGB\")\n        label = self.df.loc[idx, \"diagnosis\"]\n\n        img = self.transform(image=np.array(img))[\"image\"]\n        return img, label\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = APTOSDataset(train_df, train_tfms)\nval_ds   = APTOSDataset(val_df, val_tfms)\ntest_ds  = APTOSDataset(test_df, val_tfms)\n\ntrain_loader = DataLoader(train_ds, batch_size=16, shuffle=True, num_workers=2)\nval_loader   = DataLoader(val_ds, batch_size=16, shuffle=False)\ntest_loader  = DataLoader(test_ds, batch_size=16, shuffle=False)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CBAM(nn.Module):\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n\n        self.mlp = nn.Sequential(\n            nn.Conv2d(channels, channels // reduction, 1),\n            nn.ReLU(),\n            nn.Conv2d(channels // reduction, channels, 1)\n        )\n\n        self.spatial = nn.Conv2d(2, 1, kernel_size=7, padding=3)\n\n    def forward(self, x):\n        ca = torch.sigmoid(self.mlp(self.avg_pool(x)) + self.mlp(self.max_pool(x)))\n        x = x * ca\n\n        sa = torch.cat([x.mean(1, keepdim=True), x.max(1, keepdim=True)[0]], dim=1)\n        sa = torch.sigmoid(self.spatial(sa))\n        return x * sa\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CNN_ViT(nn.Module):\n    def __init__(self, num_classes=5):\n        super().__init__()\n\n        self.backbone = models.densenet121(pretrained=True)\n        self.backbone = nn.Sequential(*list(self.backbone.features.children()))\n        self.attn = CBAM(1024)\n\n        self.vit = timm.create_model(\n            \"vit_base_patch16_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.pool = nn.AdaptiveAvgPool2d((1,1))\n        self.fc = nn.Linear(1024 + 768, num_classes)\n\n    def forward(self, x):\n        cnn_feat = self.backbone(x)\n        cnn_feat = self.attn(cnn_feat)\n        cnn_feat = self.pool(cnn_feat).flatten(1)\n\n        vit_feat = self.vit(x)\n\n        feat = torch.cat([cnn_feat, vit_feat], dim=1)\n        return self.fc(feat)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n\n    def forward(self, logits, targets):\n        ce = F.cross_entropy(logits, targets, reduction=\"none\")\n        pt = torch.exp(-ce)\n        return ((1 - pt) ** self.gamma * ce).mean()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = CNN_ViT().to(device)\ncriterion = FocalLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(loader):\n    model.train()\n    correct, total, loss_sum = 0, 0, 0\n\n    for x, y in loader:\n        x, y = x.to(device), y.to(device)\n        optimizer.zero_grad()\n        out = model(x)\n        loss = criterion(out, y)\n        loss.backward()\n        optimizer.step()\n\n        loss_sum += loss.item()\n        correct += (out.argmax(1) == y).sum().item()\n        total += y.size(0)\n\n    return loss_sum/len(loader), correct/total\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def eval_epoch(loader):\n    model.eval()\n    correct, total = 0, 0\n\n    with torch.no_grad():\n        for x, y in loader:\n            x, y = x.to(device), y.to(device)\n            out = model(x)\n            correct += (out.argmax(1) == y).sum().item()\n            total += y.size(0)\n\n    return correct/total\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 10\n\nfor epoch in range(EPOCHS):\n    train_loss, train_acc = train_epoch(train_loader)\n    val_acc = eval_epoch(val_loader)\n\n    print(f\"Epoch {epoch+1}/{EPOCHS} | \"\n          f\"Train Acc: {train_acc:.4f} | Val Acc: {val_acc:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}