{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":4104,"databundleVersionId":46661,"sourceType":"competition"},{"sourceId":169620,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":144309,"modelId":166879}],"dockerImageVersionId":30786,"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,"execution":{"iopub.status.busy":"2024-11-17T14:10:16.810051Z","iopub.execute_input":"2024-11-17T14:10:16.810468Z","iopub.status.idle":"2024-11-17T14:10:17.933514Z","shell.execute_reply.started":"2024-11-17T14:10:16.810424Z","shell.execute_reply":"2024-11-17T14:10:17.932297Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader\nfrom torchvision.datasets import ImageFolder\nfrom tqdm import tqdm\nimport os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:17.935792Z","iopub.execute_input":"2024-11-17T14:10:17.936426Z","iopub.status.idle":"2024-11-17T14:10:22.714379Z","shell.execute_reply.started":"2024-11-17T14:10:17.936384Z","shell.execute_reply":"2024-11-17T14:10:22.713178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom functools import partial\n\nfrom timm.models.layers import DropPath, trunc_normal_\n\nclass Mlp(nn.Module):\n    def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):\n        super().__init__()\n        out_features = out_features or in_features\n        hidden_features = hidden_features or in_features\n        self.fc1 = nn.Linear(in_features, hidden_features)\n        self.act = act_layer()\n        self.fc2 = nn.Linear(hidden_features, out_features)\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.act(x)\n        x = self.drop(x)\n        x = self.fc2(x)\n        x = self.drop(x)\n        return x\n\n\nclass Attention(nn.Module):\n    def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0.):\n        super().__init__()\n        self.num_heads = num_heads\n        head_dim = dim // num_heads\n        # NOTE scale factor was wrong in my original version, can set manually to be compat with prev weights\n        self.scale = qk_scale or head_dim ** -0.5\n\n        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n    def forward(self, x):\n        B, N, C = x.shape\n        qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]  # make torchscript happy (cannot use tensor as tuple)\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale\n        attn = attn.softmax(dim=-1)\n        attn = self.attn_drop(attn)\n\n        x = (attn @ v).transpose(1, 2).reshape(B, N, C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n\n\nclass Block(nn.Module):\n\n    def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,\n                 drop_path=0., act_layer=nn.GELU, norm_layer=partial(nn.LayerNorm, eps=1e-6)):\n        super().__init__()\n        self.norm1 = norm_layer(dim)\n        self.attn = Attention(\n            dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)\n        # NOTE: drop path for stochastic depth, we shall see if this is better than dropout here\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n        self.norm2 = norm_layer(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n\n    def forward(self, x):\n        x = x + self.drop_path(self.attn(self.norm1(x)))\n        x = x + self.drop_path(self.mlp(self.norm2(x)))\n        return x\n\n\nclass ConvBlock(nn.Module):\n\n    def __init__(self, inplanes, outplanes, stride=1, res_conv=False, act_layer=nn.ReLU, groups=1,\n                 norm_layer=partial(nn.BatchNorm2d, eps=1e-6), drop_block=None, drop_path=None):\n        super(ConvBlock, self).__init__()\n\n        expansion = 4\n        med_planes = outplanes // expansion\n\n        self.conv1 = nn.Conv2d(inplanes, med_planes, kernel_size=1, stride=1, padding=0, bias=False)\n        self.bn1 = norm_layer(med_planes)\n        self.act1 = act_layer(inplace=True)\n\n        self.conv2 = nn.Conv2d(med_planes, med_planes, kernel_size=3, stride=stride, groups=groups, padding=1, bias=False)\n        self.bn2 = norm_layer(med_planes)\n        self.act2 = act_layer(inplace=True)\n\n        self.conv3 = nn.Conv2d(med_planes, outplanes, kernel_size=1, stride=1, padding=0, bias=False)\n        self.bn3 = norm_layer(outplanes)\n        self.act3 = act_layer(inplace=True)\n\n        if res_conv:\n            self.residual_conv = nn.Conv2d(inplanes, outplanes, kernel_size=1, stride=stride, padding=0, bias=False)\n            self.residual_bn = norm_layer(outplanes)\n\n        self.res_conv = res_conv\n        self.drop_block = drop_block\n        self.drop_path = drop_path\n\n    def zero_init_last_bn(self):\n        nn.init.zeros_(self.bn3.weight)\n\n    def forward(self, x, x_t=None, return_x_2=True):\n        residual = x\n\n        x = self.conv1(x)\n        x = self.bn1(x)\n        if self.drop_block is not None:\n            x = self.drop_block(x)\n        x = self.act1(x)\n\n        x = self.conv2(x) if x_t is None else self.conv2(x + x_t)\n        x = self.bn2(x)\n        if self.drop_block is not None:\n            x = self.drop_block(x)\n        x2 = self.act2(x)\n\n        x = self.conv3(x2)\n        x = self.bn3(x)\n        if self.drop_block is not None:\n            x = self.drop_block(x)\n\n        if self.drop_path is not None:\n            x = self.drop_path(x)\n\n        if self.res_conv:\n            residual = self.residual_conv(residual)\n            residual = self.residual_bn(residual)\n\n        x += residual\n        x = self.act3(x)\n\n        if return_x_2:\n            return x, x2\n        else:\n            return x\n\n\nclass FCUDown(nn.Module):\n    \"\"\" CNN feature maps -> Transformer patch embeddings\n    \"\"\"\n\n    def __init__(self, inplanes, outplanes, dw_stride, act_layer=nn.GELU,\n                 norm_layer=partial(nn.LayerNorm, eps=1e-6)):\n        super(FCUDown, self).__init__()\n        self.dw_stride = dw_stride\n\n        self.conv_project = nn.Conv2d(inplanes, outplanes, kernel_size=1, stride=1, padding=0)\n        self.sample_pooling = nn.AvgPool2d(kernel_size=dw_stride, stride=dw_stride)\n\n        self.ln = norm_layer(outplanes)\n        self.act = act_layer()\n\n    def forward(self, x, x_t):\n        x = self.conv_project(x)  # [N, C, H, W]\n\n        x = self.sample_pooling(x).flatten(2).transpose(1, 2)\n        x = self.ln(x)\n        x = self.act(x)\n\n        x = torch.cat([x_t[:, 0][:, None, :], x], dim=1)\n\n        return x\n\n\nclass FCUUp(nn.Module):\n    \"\"\" Transformer patch embeddings -> CNN feature maps\n    \"\"\"\n\n    def __init__(self, inplanes, outplanes, up_stride, act_layer=nn.ReLU,\n                 norm_layer=partial(nn.BatchNorm2d, eps=1e-6),):\n        super(FCUUp, self).__init__()\n\n        self.up_stride = up_stride\n        self.conv_project = nn.Conv2d(inplanes, outplanes, kernel_size=1, stride=1, padding=0)\n        self.bn = norm_layer(outplanes)\n        self.act = act_layer()\n\n    def forward(self, x, H, W):\n        B, _, C = x.shape\n        # [N, 197, 384] -> [N, 196, 384] -> [N, 384, 196] -> [N, 384, 14, 14]\n        x_r = x[:, 1:].transpose(1, 2).reshape(B, C, H, W)\n        x_r = self.act(self.bn(self.conv_project(x_r)))\n\n        return F.interpolate(x_r, size=(H * self.up_stride, W * self.up_stride))\n\n\nclass Med_ConvBlock(nn.Module):\n    \"\"\" special case for Convblock with down sampling,\n    \"\"\"\n    def __init__(self, inplanes, act_layer=nn.ReLU, groups=1, norm_layer=partial(nn.BatchNorm2d, eps=1e-6),\n                 drop_block=None, drop_path=None):\n\n        super(Med_ConvBlock, self).__init__()\n\n        expansion = 4\n        med_planes = inplanes // expansion\n\n        self.conv1 = nn.Conv2d(inplanes, med_planes, kernel_size=1, stride=1, padding=0, bias=False)\n        self.bn1 = norm_layer(med_planes)\n        self.act1 = act_layer(inplace=True)\n\n        self.conv2 = nn.Conv2d(med_planes, med_planes, kernel_size=3, stride=1, groups=groups, padding=1, bias=False)\n        self.bn2 = norm_layer(med_planes)\n        self.act2 = act_layer(inplace=True)\n\n        self.conv3 = nn.Conv2d(med_planes, inplanes, kernel_size=1, stride=1, padding=0, bias=False)\n        self.bn3 = norm_layer(inplanes)\n        self.act3 = act_layer(inplace=True)\n\n        self.drop_block = drop_block\n        self.drop_path = drop_path\n\n    def zero_init_last_bn(self):\n        nn.init.zeros_(self.bn3.weight)\n\n    def forward(self, x):\n        residual = x\n\n        x = self.conv1(x)\n        x = self.bn1(x)\n        if self.drop_block is not None:\n            x = self.drop_block(x)\n        x = self.act1(x)\n\n        x = self.conv2(x)\n        x = self.bn2(x)\n        if self.drop_block is not None:\n            x = self.drop_block(x)\n        x = self.act2(x)\n\n        x = self.conv3(x)\n        x = self.bn3(x)\n        if self.drop_block is not None:\n            x = self.drop_block(x)\n\n        if self.drop_path is not None:\n            x = self.drop_path(x)\n\n        x += residual\n        x = self.act3(x)\n\n        return x\n\n\nclass ConvTransBlock(nn.Module):\n    \"\"\"\n    Basic module for ConvTransformer, keep feature maps for CNN block and patch embeddings for transformer encoder block\n    \"\"\"\n\n    def __init__(self, inplanes, outplanes, res_conv, stride, dw_stride, embed_dim, num_heads=12, mlp_ratio=4.,\n                 qkv_bias=False, qk_scale=None, drop_rate=0., attn_drop_rate=0., drop_path_rate=0.,\n                 last_fusion=False, num_med_block=0, groups=1):\n\n        super(ConvTransBlock, self).__init__()\n        expansion = 4\n        self.cnn_block = ConvBlock(inplanes=inplanes, outplanes=outplanes, res_conv=res_conv, stride=stride, groups=groups)\n\n        if last_fusion:\n            self.fusion_block = ConvBlock(inplanes=outplanes, outplanes=outplanes, stride=2, res_conv=True, groups=groups)\n        else:\n            self.fusion_block = ConvBlock(inplanes=outplanes, outplanes=outplanes, groups=groups)\n\n        if num_med_block > 0:\n            self.med_block = []\n            for i in range(num_med_block):\n                self.med_block.append(Med_ConvBlock(inplanes=outplanes, groups=groups))\n            self.med_block = nn.ModuleList(self.med_block)\n\n        self.squeeze_block = FCUDown(inplanes=outplanes // expansion, outplanes=embed_dim, dw_stride=dw_stride)\n\n        self.expand_block = FCUUp(inplanes=embed_dim, outplanes=outplanes // expansion, up_stride=dw_stride)\n\n        self.trans_block = Block(\n            dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,\n            drop=drop_rate, attn_drop=attn_drop_rate, drop_path=drop_path_rate)\n\n        self.dw_stride = dw_stride\n        self.embed_dim = embed_dim\n        self.num_med_block = num_med_block\n        self.last_fusion = last_fusion\n\n    def forward(self, x, x_t):\n        x, x2 = self.cnn_block(x)\n\n        _, _, H, W = x2.shape\n\n        x_st = self.squeeze_block(x2, x_t)\n\n        x_t = self.trans_block(x_st + x_t)\n\n        if self.num_med_block > 0:\n            for m in self.med_block:\n                x = m(x)\n\n        x_t_r = self.expand_block(x_t, H // self.dw_stride, W // self.dw_stride)\n        x = self.fusion_block(x, x_t_r, return_x_2=False)\n\n        return x, x_t\n\n\nclass Conformer(nn.Module):\n\n    def __init__(self, patch_size=16, in_chans=3, num_classes=1000, base_channel=64, channel_ratio=4, num_med_block=0,\n                 embed_dim=768, depth=12, num_heads=12, mlp_ratio=4., qkv_bias=False, qk_scale=None,\n                 drop_rate=0., attn_drop_rate=0., drop_path_rate=0.):\n\n        # Transformer\n        super().__init__()\n        self.num_classes = num_classes\n        self.num_features = self.embed_dim = embed_dim  # num_features for consistency with other models\n        assert depth % 3 == 0\n\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))\n        self.trans_dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]  # stochastic depth decay rule\n\n        # Classifier head\n        self.trans_norm = nn.LayerNorm(embed_dim)\n        self.trans_cls_head = nn.Linear(embed_dim, num_classes) if num_classes > 0 else nn.Identity()\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.conv_cls_head = nn.Linear(int(256 * channel_ratio), num_classes)\n\n        # Stem stage: get the feature maps by conv block (copied form resnet.py)\n        self.conv1 = nn.Conv2d(in_chans, 64, kernel_size=7, stride=2, padding=3, bias=False)  # 1 / 2 [112, 112]\n        self.bn1 = nn.BatchNorm2d(64)\n        self.act1 = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)  # 1 / 4 [56, 56]\n\n        # 1 stage\n        stage_1_channel = int(base_channel * channel_ratio)\n        trans_dw_stride = patch_size // 4\n        self.conv_1 = ConvBlock(inplanes=64, outplanes=stage_1_channel, res_conv=True, stride=1)\n        self.trans_patch_conv = nn.Conv2d(64, embed_dim, kernel_size=trans_dw_stride, stride=trans_dw_stride, padding=0)\n        self.trans_1 = Block(dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias,\n                             qk_scale=qk_scale, drop=drop_rate, attn_drop=attn_drop_rate, drop_path=self.trans_dpr[0],\n                             )\n\n        # 2~4 stage\n        init_stage = 2\n        fin_stage = depth // 3 + 1\n        for i in range(init_stage, fin_stage):\n            self.add_module('conv_trans_' + str(i),\n                    ConvTransBlock(\n                        stage_1_channel, stage_1_channel, False, 1, dw_stride=trans_dw_stride, embed_dim=embed_dim,\n                        num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,\n                        drop_rate=drop_rate, attn_drop_rate=attn_drop_rate, drop_path_rate=self.trans_dpr[i-1],\n                        num_med_block=num_med_block\n                    )\n            )\n\n\n        stage_2_channel = int(base_channel * channel_ratio * 2)\n        # 5~8 stage\n        init_stage = fin_stage # 5\n        fin_stage = fin_stage + depth // 3 # 9\n        for i in range(init_stage, fin_stage):\n            s = 2 if i == init_stage else 1\n            in_channel = stage_1_channel if i == init_stage else stage_2_channel\n            res_conv = True if i == init_stage else False\n            self.add_module('conv_trans_' + str(i),\n                    ConvTransBlock(\n                        in_channel, stage_2_channel, res_conv, s, dw_stride=trans_dw_stride // 2, embed_dim=embed_dim,\n                        num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,\n                        drop_rate=drop_rate, attn_drop_rate=attn_drop_rate, drop_path_rate=self.trans_dpr[i-1],\n                        num_med_block=num_med_block\n                    )\n            )\n\n        stage_3_channel = int(base_channel * channel_ratio * 2 * 2)\n        # 9~12 stage\n        init_stage = fin_stage  # 9\n        fin_stage = fin_stage + depth // 3  # 13\n        for i in range(init_stage, fin_stage):\n            s = 2 if i == init_stage else 1\n            in_channel = stage_2_channel if i == init_stage else stage_3_channel\n            res_conv = True if i == init_stage else False\n            last_fusion = True if i == depth else False\n            self.add_module('conv_trans_' + str(i),\n                    ConvTransBlock(\n                        in_channel, stage_3_channel, res_conv, s, dw_stride=trans_dw_stride // 4, embed_dim=embed_dim,\n                        num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,\n                        drop_rate=drop_rate, attn_drop_rate=attn_drop_rate, drop_path_rate=self.trans_dpr[i-1],\n                        num_med_block=num_med_block, last_fusion=last_fusion\n                    )\n            )\n        self.fin_stage = fin_stage\n\n        trunc_normal_(self.cls_token, std=.02)\n\n        self.apply(self._init_weights)\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Linear):\n            trunc_normal_(m.weight, std=.02)\n            if isinstance(m, nn.Linear) and m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.bias, 0)\n            nn.init.constant_(m.weight, 1.0)\n        elif isinstance(m, nn.Conv2d):\n            nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n        elif isinstance(m, nn.BatchNorm2d):\n            nn.init.constant_(m.weight, 1.)\n            nn.init.constant_(m.bias, 0.)\n        elif isinstance(m, nn.GroupNorm):\n            nn.init.constant_(m.weight, 1.)\n            nn.init.constant_(m.bias, 0.)\n\n    @torch.jit.ignore\n    def no_weight_decay(self):\n        return {'cls_token'}\n\n\n    def forward(self, x):\n        B = x.shape[0]\n        cls_tokens = self.cls_token.expand(B, -1, -1)\n\n        # pdb.set_trace()\n        # stem stage [N, 3, 224, 224] -> [N, 64, 56, 56]\n        x_base = self.maxpool(self.act1(self.bn1(self.conv1(x))))\n\n        # 1 stage\n        x = self.conv_1(x_base, return_x_2=False)\n\n        x_t = self.trans_patch_conv(x_base).flatten(2).transpose(1, 2)\n        x_t = torch.cat([cls_tokens, x_t], dim=1)\n        x_t = self.trans_1(x_t)\n        \n        # 2 ~ final \n        for i in range(2, self.fin_stage):\n            x, x_t = eval('self.conv_trans_' + str(i))(x, x_t)\n\n        # conv classification\n        x_p = self.pooling(x).flatten(1)\n        conv_cls = self.conv_cls_head(x_p)\n\n        # trans classification\n        x_t = self.trans_norm(x_t)\n        tran_cls = self.trans_cls_head(x_t[:, 0])\n\n        return [conv_cls, tran_cls]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:22.716673Z","iopub.execute_input":"2024-11-17T14:10:22.717338Z","iopub.status.idle":"2024-11-17T14:10:24.751894Z","shell.execute_reply.started":"2024-11-17T14:10:22.717285Z","shell.execute_reply":"2024-11-17T14:10:24.750873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:24.754460Z","iopub.execute_input":"2024-11-17T14:10:24.754829Z","iopub.status.idle":"2024-11-17T14:10:24.760001Z","shell.execute_reply.started":"2024-11-17T14:10:24.754790Z","shell.execute_reply":"2024-11-17T14:10:24.758883Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch_size = 16\nnum_epochs = 10\nlearning_rate = 1e-4\nweight_decay = 1e-5\ndrop_rate = 0.1\nattn_drop_rate = 0.1\ndrop_path_rate = 0.1\npatch_size = 16\nnum_classes = 1000 \nembed_dim = 768\nnum_heads = 12\nmlp_ratio = 4.0\nqkv_bias = False\nqk_scale = None\nnum_med_block = 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:24.761370Z","iopub.execute_input":"2024-11-17T14:10:24.761725Z","iopub.status.idle":"2024-11-17T14:10:24.773468Z","shell.execute_reply.started":"2024-11-17T14:10:24.761688Z","shell.execute_reply":"2024-11-17T14:10:24.772392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!unzip ../input/diabetic-retinopathy-detection/trainLabels.csv.zip\n!apt install p7zip-full -y\n!7z x ../input/diabetic-retinopathy-detection/train.zip.001 \"-i!train/11*.jpeg\" -y # restrict extracted file to about 100 for the disk restriction\n!mkdir data\n!mv train data/train_11","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:19:24.434380Z","iopub.execute_input":"2024-11-17T14:19:24.434854Z","iopub.status.idle":"2024-11-17T14:19:53.490339Z","shell.execute_reply.started":"2024-11-17T14:19:24.434804Z","shell.execute_reply":"2024-11-17T14:19:53.488867Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\n\n# Define paths and read the CSV\nbase_image_dir = os.path.join('/', 'kaggle', 'working', 'data', 'train_11')\nretina_df = pd.read_csv('/kaggle/working/trainLabels.csv')\n\n# Add columns for paths and check existence\nretina_df['PatientId'] = retina_df['image'].map(lambda x: x.split('_')[0])\nretina_df['path'] = retina_df['image'].map(lambda x: os.path.join(base_image_dir, f\"{x}.jpeg\"))\nretina_df['exists'] = retina_df['path'].map(os.path.exists)\n\n# Filter data\nretina_df.dropna(inplace=True)\nretina_df = retina_df[retina_df['exists']]\n\n# Split into train and validation sets\nrr_df = retina_df[['PatientId', 'level']].drop_duplicates()\ntrain_ids, valid_ids = train_test_split(\n    rr_df['PatientId'],\n    test_size=0.25,\n    random_state=2018,\n    stratify=rr_df['level']\n)\ntrain_df = retina_df[retina_df['PatientId'].isin(train_ids)]\nvalid_df = retina_df[retina_df['PatientId'].isin(valid_ids)]\n\nprint(f\"Train size: {train_df.shape[0]}, Validation size: {valid_df.shape[0]}\")\n\n# Custom Dataset Class\nclass RetinaDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        row = self.dataframe.iloc[idx]\n        image = Image.open(row['path']).convert('RGB')  # Load the image\n        label = row['level']  # Load the label\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\n# Define transforms for data augmentation and preprocessing\ntrain_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),  # Resize images to a fixed size\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ToTensor()\n])\n\nvalid_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\n# Create Dataset and DataLoader\ntrain_dataset = RetinaDataset(train_df, transform=train_transforms)\nvalid_dataset = RetinaDataset(valid_df, transform=valid_transforms)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=32, shuffle=False)\n\n# Check a sample batch\nfor images, labels in train_loader:\n    print(f\"Images batch shape: {images.size()}\")\n    print(f\"Labels batch shape: {labels.size()}\")\n    break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:20:08.629128Z","iopub.execute_input":"2024-11-17T14:20:08.629622Z","iopub.status.idle":"2024-11-17T14:20:14.976302Z","shell.execute_reply.started":"2024-11-17T14:20:08.629578Z","shell.execute_reply":"2024-11-17T14:20:14.974900Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = Conformer(\n    patch_size=patch_size,\n    in_chans=3,\n    num_classes=num_classes,\n    base_channel=64,\n    channel_ratio=4,\n    num_med_block=num_med_block,\n    embed_dim=embed_dim,\n    depth=12,\n    num_heads=num_heads,\n    mlp_ratio=mlp_ratio,\n    qkv_bias=qkv_bias,\n    qk_scale=qk_scale,\n    drop_rate=drop_rate,\n    attn_drop_rate=attn_drop_rate,\n    drop_path_rate=drop_path_rate\n).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:22:37.779034Z","iopub.execute_input":"2024-11-17T14:22:37.779483Z","iopub.status.idle":"2024-11-17T14:22:40.080345Z","shell.execute_reply.started":"2024-11-17T14:22:37.779436Z","shell.execute_reply":"2024-11-17T14:22:40.079234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pretrained_path = '/kaggle/input/conformer_base_patch_16/pytorch/default/1/Conformer_base_patch16.pth'  # Replace with your .pth file path\ncheckpoint = torch.load(pretrained_path, map_location=device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:23:09.768650Z","iopub.execute_input":"2024-11-17T14:23:09.769151Z","iopub.status.idle":"2024-11-17T14:23:12.108740Z","shell.execute_reply.started":"2024-11-17T14:23:09.769104Z","shell.execute_reply":"2024-11-17T14:23:12.107510Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(checkpoint)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:23:21.541394Z","iopub.execute_input":"2024-11-17T14:23:21.542400Z","iopub.status.idle":"2024-11-17T14:23:21.827892Z","shell.execute_reply.started":"2024-11-17T14:23:21.542354Z","shell.execute_reply":"2024-11-17T14:23:21.824037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:25.663409Z","iopub.status.idle":"2024-11-17T14:10:25.663979Z","shell.execute_reply.started":"2024-11-17T14:10:25.663674Z","shell.execute_reply":"2024-11-17T14:10:25.663702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:25.665705Z","iopub.status.idle":"2024-11-17T14:10:25.666272Z","shell.execute_reply.started":"2024-11-17T14:10:25.665979Z","shell.execute_reply":"2024-11-17T14:10:25.666009Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(epoch, model, train_loader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    correct_preds = 0\n    total_preds = 0\n    for inputs, labels in tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs}'):\n        inputs, labels = inputs.to(device), labels.to(device)\n        \n        # Zero the gradients\n        optimizer.zero_grad()\n        \n        # Forward pass\n        outputs = model(inputs)\n        \n        # Calculate loss\n        loss = criterion(outputs[0], labels)  # Use the first classifier (assuming it's conv_cls)\n        \n        # Backward pass\n        loss.backward()\n        optimizer.step()\n        \n        # Track loss and accuracy\n        running_loss += loss.item()\n        _, predicted = outputs[0].max(1)\n        correct_preds += (predicted == labels).sum().item()\n        total_preds += labels.size(0)\n\n    epoch_loss = running_loss / len(train_loader)\n    epoch_accuracy = 100. * correct_preds / total_preds\n    return epoch_loss, epoch_accuracy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:25.667480Z","iopub.status.idle":"2024-11-17T14:10:25.668053Z","shell.execute_reply.started":"2024-11-17T14:10:25.667742Z","shell.execute_reply":"2024-11-17T14:10:25.667771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, val_loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct_preds = 0\n    total_preds = 0\n    with torch.no_grad():\n        for inputs, labels in tqdm(val_loader, desc='Validation'):\n            inputs, labels = inputs.to(device), labels.to(device)\n\n            outputs = model(inputs)\n\n            loss = criterion(outputs[0], labels)\n\n            running_loss += loss.item()\n            _, predicted = outputs[0].max(1)\n            correct_preds += (predicted == labels).sum().item()\n            total_preds += labels.size(0)\n\n    val_loss = running_loss / len(val_loader)\n    val_accuracy = 100. * correct_preds / total_preds\n    return val_loss, val_accuracy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:25.670227Z","iopub.status.idle":"2024-11-17T14:10:25.670645Z","shell.execute_reply.started":"2024-11-17T14:10:25.670432Z","shell.execute_reply":"2024-11-17T14:10:25.670453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_val_accuracy = 0.0\nfor epoch in range(num_epochs):\n    # Train for one epoch\n    train_loss, train_accuracy = train_one_epoch(epoch, model, train_loader, optimizer, criterion, device)\n    \n    # Validate\n    val_loss, val_accuracy = validate(model, val_loader, criterion, device)\n    \n    print(f\"Epoch {epoch+1}/{num_epochs}\")\n    print(f\"Train Loss: {train_loss:.4f}, Train Accuracy: {train_accuracy:.2f}%\")\n    print(f\"Validation Loss: {val_loss:.4f}, Validation Accuracy: {val_accuracy:.2f}%\")\n    \n    # Save the best model based on validation accuracy\n    if val_accuracy > best_val_accuracy:\n        best_val_accuracy = val_accuracy\n        torch.save(model.state_dict(), 'best_conformer_model.pth')\n\n    # Step the learning rate scheduler\n    scheduler.step()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-17T14:10:25.672572Z","iopub.status.idle":"2024-11-17T14:10:25.673037Z","shell.execute_reply.started":"2024-11-17T14:10:25.672797Z","shell.execute_reply":"2024-11-17T14:10:25.672818Z"}},"outputs":[],"execution_count":null}]}