{"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","modelInstanceId":144309,"modelId":166879},{"sourceId":170827,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":145361,"modelId":167922}],"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-18T16:06:53.288606Z","iopub.execute_input":"2024-11-18T16:06:53.289399Z","iopub.status.idle":"2024-11-18T16:06:53.300234Z","shell.execute_reply.started":"2024-11-18T16:06:53.289352Z","shell.execute_reply":"2024-11-18T16:06:53.299025Z"}},"outputs":[{"name":"stdout","text":"/kaggle/input/conformer/other/default/1/Conformer_small_patch16.pth\n/kaggle/input/diabetic-retinopathy-detection/train.zip.003\n/kaggle/input/diabetic-retinopathy-detection/test.zip.004\n/kaggle/input/diabetic-retinopathy-detection/test.zip.005\n/kaggle/input/diabetic-retinopathy-detection/train.zip.002\n/kaggle/input/diabetic-retinopathy-detection/test.zip.006\n/kaggle/input/diabetic-retinopathy-detection/test.zip.003\n/kaggle/input/diabetic-retinopathy-detection/train.zip.005\n/kaggle/input/diabetic-retinopathy-detection/train.zip.001\n/kaggle/input/diabetic-retinopathy-detection/sampleSubmission.csv.zip\n/kaggle/input/diabetic-retinopathy-detection/test.zip.007\n/kaggle/input/diabetic-retinopathy-detection/trainLabels.csv.zip\n/kaggle/input/diabetic-retinopathy-detection/test.zip.001\n/kaggle/input/diabetic-retinopathy-detection/sample.zip\n/kaggle/input/diabetic-retinopathy-detection/train.zip.004\n/kaggle/input/diabetic-retinopathy-detection/test.zip.002\n","output_type":"stream"}],"execution_count":10},{"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-18T16:06:53.568204Z","iopub.execute_input":"2024-11-18T16:06:53.568686Z","iopub.status.idle":"2024-11-18T16:06:53.575164Z","shell.execute_reply.started":"2024-11-18T16:06:53.568641Z","shell.execute_reply":"2024-11-18T16:06:53.573899Z"}},"outputs":[],"execution_count":11},{"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-18T16:06:54.216711Z","iopub.execute_input":"2024-11-18T16:06:54.217142Z","iopub.status.idle":"2024-11-18T16:06:54.314874Z","shell.execute_reply.started":"2024-11-18T16:06:54.217102Z","shell.execute_reply":"2024-11-18T16:06:54.313656Z"}},"outputs":[],"execution_count":12},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T16:06:54.316917Z","iopub.execute_input":"2024-11-18T16:06:54.317404Z","iopub.status.idle":"2024-11-18T16:06:54.330559Z","shell.execute_reply.started":"2024-11-18T16:06:54.317351Z","shell.execute_reply":"2024-11-18T16:06:54.329387Z"}},"outputs":[],"execution_count":13},{"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-18T16:06:54.633086Z","iopub.execute_input":"2024-11-18T16:06:54.633541Z","iopub.status.idle":"2024-11-18T16:06:54.641006Z","shell.execute_reply.started":"2024-11-18T16:06:54.633493Z","shell.execute_reply":"2024-11-18T16:06:54.639662Z"}},"outputs":[],"execution_count":14},{"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-18T16:06:54.950156Z","iopub.execute_input":"2024-11-18T16:06:54.950614Z","iopub.status.idle":"2024-11-18T16:10:10.759355Z","shell.execute_reply.started":"2024-11-18T16:06:54.950571Z","shell.execute_reply":"2024-11-18T16:10:10.757758Z"}},"outputs":[{"name":"stdout","text":"Archive:  ../input/diabetic-retinopathy-detection/trainLabels.csv.zip\nreplace trainLabels.csv? [y]es, [n]o, [A]ll, [N]one, [r]ename: ^C\nReading package lists... Done\nBuilding dependency tree       \nReading state information... Done\np7zip-full is already the newest version (16.02+dfsg-7build1).\n0 upgraded, 0 newly installed, 0 to remove and 46 not upgraded.\n\n7-Zip [64] 16.02 : Copyright (c) 1999-2016 Igor Pavlov : 2016-05-21\np7zip Version 16.02 (locale=C,Utf16=off,HugeFiles=on,64 bits,4 CPUs Intel(R) Xeon(R) CPU @ 2.20GHz (406F0),ASM,AES-NI)\n\nScanning the drive for archives:\n  0M Scan ../input/diabetic-retinopathy-detectio                                                1 file, 8388608000 bytes (8000 MiB)\n\nExtracting archive: ../input/diabetic-retinopathy-detection/train.zip.001\n  0% 1 Ope          --\nPath = ../input/diabetic-retinopathy-detection/train.zip.001\nType = Split\nPhysical Size = 8388608000\nVolumes = 5\nTotal Physical Size = 34988445506\n----\nPath = train.zip\nSize = 34988445506\n--\nPath = train.zip\nType = zip\nPhysical Size = 34988445506\n64-bit = +\n\n      1% 10 - train/11010_left.jpe                                2% 22 - train/11032_left.jpe                                4% 34 - train/11048_left.jpe                                6% 49 - train/11058_right.jp                                7% 63 - train/11083_right.jp                                9% 76 - train/11098_left.jpe                               10% 93 - train/11119_right.jp                               12% 110 - train/11129_left.jp                               14% 124 - train/11143_left.jp                               15% 136 - train/11156_left.jp                               17% 153 - train/11165_right.jpe                                 18% 168 - train/11183_left.jp                               20% 187 - train/11203_right.jpe                                 21% 202 - train/11217_left.jp                               23% 215 - train/11230_right.jpe                                 25% 226 - train/11241_left.jp                               26% 238 - train/11261_left.jp                               28% 252 - train/11271_left.jp                               30% 262 - train/11289_left.jp                               31% 276 - train/11301_left.jp                               33% 286 - train/11313_left.jp                               34% 297 - train/11319_right.jpe                                 36% 309 - train/11335_right.jpe                                 37% 326 - train/11364_left.jp                               39% 340 - train/11378_left.jp                               40% 350 - train/11388_left.jp                               42% 364 - train/11397_left.jp                               44% 376 - train/11417_left.jp                               45% 389 - train/1142_right.jp                               47% 408 - train/11446_left.jp                               48% 421 - train/11464_right.jpe                                 50% 435 - train/11484_right.jpe                                 51% 449 - train/11496_right.jpe                                 53% 459 - train/11505_right.jpe                                 54% 475 - train/1152_right.jp                               56% 490 - train/11547_left.jp                               57% 505 - train/11575_right.jpe                                 59% 521 - train/115_right.jpe                               61% 537 - train/11616_right.jpe                                 62% 555 - train/11630_right.jpe                                 64% 567 - train/11645_right.jpe                                 65% 582 - train/11655_left.jp                               67% 595 - train/1167_right.jp                               68% 611 - train/1170_right.jp                               70% 626 - train/11730_left.jp                               72% 638 - train/11744_left.jp                               73% 652 - train/11768_left.jp                               75% 665 - train/11778_right.jpe                                 76% 677 - train/11789_right.jpe                                 78% 687 - train/11796_right.jpe                                 79% 700 - train/1180_left.jpe                               81% 711 - train/11818_right.jpe                                 82% 725 - train/1182_right.jp                               84% 742 - train/11854_left.jp                               85% 754 - train/11870_left.jp                               87% 770 - train/11896_left.jp                               89% 781 - train/11909_right.jpe                                 90% 793 - train/11920_right.jpe                                 92% 806 - train/11937_left.jp                               93% 818 - train/11952_left.jp                               95% 832 - train/11966_left.jp                               97% 844 - train/11975_left.jp                               98% 859 - train/11990_right.jpe                                Everything is Ok\n\nFiles: 872\nSize:       945260232\nCompressed: 34988445506\nmkdir: cannot create directory 'data': File exists\n","output_type":"stream"}],"execution_count":15},{"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-18T16:10:28.46039Z","iopub.execute_input":"2024-11-18T16:10:28.460855Z","iopub.status.idle":"2024-11-18T16:10:33.34994Z","shell.execute_reply.started":"2024-11-18T16:10:28.460812Z","shell.execute_reply":"2024-11-18T16:10:33.348526Z"}},"outputs":[{"name":"stdout","text":"Train size: 678, Validation size: 242\nImages batch shape: torch.Size([32, 3, 224, 224])\nLabels batch shape: torch.Size([32])\n","output_type":"stream"}],"execution_count":19},{"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-18T16:10:33.352577Z","iopub.execute_input":"2024-11-18T16:10:33.353261Z","iopub.status.idle":"2024-11-18T16:10:35.632231Z","shell.execute_reply.started":"2024-11-18T16:10:33.35319Z","shell.execute_reply":"2024-11-18T16:10:35.631226Z"}},"outputs":[],"execution_count":20},{"cell_type":"code","source":"model_state_dict = model.state_dict()\nfiltered_checkpoint = {k: v for k, v in checkpoint.items() if k in model_state_dict and model_state_dict[k].shape == v.shape}\nmodel.load_state_dict(filtered_checkpoint, strict=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T16:16:30.763004Z","iopub.execute_input":"2024-11-18T16:16:30.763455Z","iopub.status.idle":"2024-11-18T16:16:30.81367Z","shell.execute_reply.started":"2024-11-18T16:16:30.763415Z","shell.execute_reply":"2024-11-18T16:16:30.812518Z"}},"outputs":[{"execution_count":26,"output_type":"execute_result","data":{"text/plain":"_IncompatibleKeys(missing_keys=['cls_token', 'trans_norm.weight', 'trans_norm.bias', 'trans_cls_head.weight', 'trans_patch_conv.weight', 'trans_patch_conv.bias', 'trans_1.norm1.weight', 'trans_1.norm1.bias', 'trans_1.attn.qkv.weight', 'trans_1.attn.proj.weight', 'trans_1.attn.proj.bias', 'trans_1.norm2.weight', 'trans_1.norm2.bias', 'trans_1.mlp.fc1.weight', 'trans_1.mlp.fc1.bias', 'trans_1.mlp.fc2.weight', 'trans_1.mlp.fc2.bias', 'conv_trans_2.squeeze_block.conv_project.weight', 'conv_trans_2.squeeze_block.conv_project.bias', 'conv_trans_2.squeeze_block.ln.weight', 'conv_trans_2.squeeze_block.ln.bias', 'conv_trans_2.expand_block.conv_project.weight', 'conv_trans_2.trans_block.norm1.weight', 'conv_trans_2.trans_block.norm1.bias', 'conv_trans_2.trans_block.attn.qkv.weight', 'conv_trans_2.trans_block.attn.proj.weight', 'conv_trans_2.trans_block.attn.proj.bias', 'conv_trans_2.trans_block.norm2.weight', 'conv_trans_2.trans_block.norm2.bias', 'conv_trans_2.trans_block.mlp.fc1.weight', 'conv_trans_2.trans_block.mlp.fc1.bias', 'conv_trans_2.trans_block.mlp.fc2.weight', 'conv_trans_2.trans_block.mlp.fc2.bias', 'conv_trans_3.squeeze_block.conv_project.weight', 'conv_trans_3.squeeze_block.conv_project.bias', 'conv_trans_3.squeeze_block.ln.weight', 'conv_trans_3.squeeze_block.ln.bias', 'conv_trans_3.expand_block.conv_project.weight', 'conv_trans_3.trans_block.norm1.weight', 'conv_trans_3.trans_block.norm1.bias', 'conv_trans_3.trans_block.attn.qkv.weight', 'conv_trans_3.trans_block.attn.proj.weight', 'conv_trans_3.trans_block.attn.proj.bias', 'conv_trans_3.trans_block.norm2.weight', 'conv_trans_3.trans_block.norm2.bias', 'conv_trans_3.trans_block.mlp.fc1.weight', 'conv_trans_3.trans_block.mlp.fc1.bias', 'conv_trans_3.trans_block.mlp.fc2.weight', 'conv_trans_3.trans_block.mlp.fc2.bias', 'conv_trans_4.squeeze_block.conv_project.weight', 'conv_trans_4.squeeze_block.conv_project.bias', 'conv_trans_4.squeeze_block.ln.weight', 'conv_trans_4.squeeze_block.ln.bias', 'conv_trans_4.expand_block.conv_project.weight', 'conv_trans_4.trans_block.norm1.weight', 'conv_trans_4.trans_block.norm1.bias', 'conv_trans_4.trans_block.attn.qkv.weight', 'conv_trans_4.trans_block.attn.proj.weight', 'conv_trans_4.trans_block.attn.proj.bias', 'conv_trans_4.trans_block.norm2.weight', 'conv_trans_4.trans_block.norm2.bias', 'conv_trans_4.trans_block.mlp.fc1.weight', 'conv_trans_4.trans_block.mlp.fc1.bias', 'conv_trans_4.trans_block.mlp.fc2.weight', 'conv_trans_4.trans_block.mlp.fc2.bias', 'conv_trans_5.squeeze_block.conv_project.weight', 'conv_trans_5.squeeze_block.conv_project.bias', 'conv_trans_5.squeeze_block.ln.weight', 'conv_trans_5.squeeze_block.ln.bias', 'conv_trans_5.expand_block.conv_project.weight', 'conv_trans_5.trans_block.norm1.weight', 'conv_trans_5.trans_block.norm1.bias', 'conv_trans_5.trans_block.attn.qkv.weight', 'conv_trans_5.trans_block.attn.proj.weight', 'conv_trans_5.trans_block.attn.proj.bias', 'conv_trans_5.trans_block.norm2.weight', 'conv_trans_5.trans_block.norm2.bias', 'conv_trans_5.trans_block.mlp.fc1.weight', 'conv_trans_5.trans_block.mlp.fc1.bias', 'conv_trans_5.trans_block.mlp.fc2.weight', 'conv_trans_5.trans_block.mlp.fc2.bias', 'conv_trans_6.squeeze_block.conv_project.weight', 'conv_trans_6.squeeze_block.conv_project.bias', 'conv_trans_6.squeeze_block.ln.weight', 'conv_trans_6.squeeze_block.ln.bias', 'conv_trans_6.expand_block.conv_project.weight', 'conv_trans_6.trans_block.norm1.weight', 'conv_trans_6.trans_block.norm1.bias', 'conv_trans_6.trans_block.attn.qkv.weight', 'conv_trans_6.trans_block.attn.proj.weight', 'conv_trans_6.trans_block.attn.proj.bias', 'conv_trans_6.trans_block.norm2.weight', 'conv_trans_6.trans_block.norm2.bias', 'conv_trans_6.trans_block.mlp.fc1.weight', 'conv_trans_6.trans_block.mlp.fc1.bias', 'conv_trans_6.trans_block.mlp.fc2.weight', 'conv_trans_6.trans_block.mlp.fc2.bias', 'conv_trans_7.squeeze_block.conv_project.weight', 'conv_trans_7.squeeze_block.conv_project.bias', 'conv_trans_7.squeeze_block.ln.weight', 'conv_trans_7.squeeze_block.ln.bias', 'conv_trans_7.expand_block.conv_project.weight', 'conv_trans_7.trans_block.norm1.weight', 'conv_trans_7.trans_block.norm1.bias', 'conv_trans_7.trans_block.attn.qkv.weight', 'conv_trans_7.trans_block.attn.proj.weight', 'conv_trans_7.trans_block.attn.proj.bias', 'conv_trans_7.trans_block.norm2.weight', 'conv_trans_7.trans_block.norm2.bias', 'conv_trans_7.trans_block.mlp.fc1.weight', 'conv_trans_7.trans_block.mlp.fc1.bias', 'conv_trans_7.trans_block.mlp.fc2.weight', 'conv_trans_7.trans_block.mlp.fc2.bias', 'conv_trans_8.squeeze_block.conv_project.weight', 'conv_trans_8.squeeze_block.conv_project.bias', 'conv_trans_8.squeeze_block.ln.weight', 'conv_trans_8.squeeze_block.ln.bias', 'conv_trans_8.expand_block.conv_project.weight', 'conv_trans_8.trans_block.norm1.weight', 'conv_trans_8.trans_block.norm1.bias', 'conv_trans_8.trans_block.attn.qkv.weight', 'conv_trans_8.trans_block.attn.proj.weight', 'conv_trans_8.trans_block.attn.proj.bias', 'conv_trans_8.trans_block.norm2.weight', 'conv_trans_8.trans_block.norm2.bias', 'conv_trans_8.trans_block.mlp.fc1.weight', 'conv_trans_8.trans_block.mlp.fc1.bias', 'conv_trans_8.trans_block.mlp.fc2.weight', 'conv_trans_8.trans_block.mlp.fc2.bias', 'conv_trans_9.squeeze_block.conv_project.weight', 'conv_trans_9.squeeze_block.conv_project.bias', 'conv_trans_9.squeeze_block.ln.weight', 'conv_trans_9.squeeze_block.ln.bias', 'conv_trans_9.expand_block.conv_project.weight', 'conv_trans_9.trans_block.norm1.weight', 'conv_trans_9.trans_block.norm1.bias', 'conv_trans_9.trans_block.attn.qkv.weight', 'conv_trans_9.trans_block.attn.proj.weight', 'conv_trans_9.trans_block.attn.proj.bias', 'conv_trans_9.trans_block.norm2.weight', 'conv_trans_9.trans_block.norm2.bias', 'conv_trans_9.trans_block.mlp.fc1.weight', 'conv_trans_9.trans_block.mlp.fc1.bias', 'conv_trans_9.trans_block.mlp.fc2.weight', 'conv_trans_9.trans_block.mlp.fc2.bias', 'conv_trans_10.squeeze_block.conv_project.weight', 'conv_trans_10.squeeze_block.conv_project.bias', 'conv_trans_10.squeeze_block.ln.weight', 'conv_trans_10.squeeze_block.ln.bias', 'conv_trans_10.expand_block.conv_project.weight', 'conv_trans_10.trans_block.norm1.weight', 'conv_trans_10.trans_block.norm1.bias', 'conv_trans_10.trans_block.attn.qkv.weight', 'conv_trans_10.trans_block.attn.proj.weight', 'conv_trans_10.trans_block.attn.proj.bias', 'conv_trans_10.trans_block.norm2.weight', 'conv_trans_10.trans_block.norm2.bias', 'conv_trans_10.trans_block.mlp.fc1.weight', 'conv_trans_10.trans_block.mlp.fc1.bias', 'conv_trans_10.trans_block.mlp.fc2.weight', 'conv_trans_10.trans_block.mlp.fc2.bias', 'conv_trans_11.squeeze_block.conv_project.weight', 'conv_trans_11.squeeze_block.conv_project.bias', 'conv_trans_11.squeeze_block.ln.weight', 'conv_trans_11.squeeze_block.ln.bias', 'conv_trans_11.expand_block.conv_project.weight', 'conv_trans_11.trans_block.norm1.weight', 'conv_trans_11.trans_block.norm1.bias', 'conv_trans_11.trans_block.attn.qkv.weight', 'conv_trans_11.trans_block.attn.proj.weight', 'conv_trans_11.trans_block.attn.proj.bias', 'conv_trans_11.trans_block.norm2.weight', 'conv_trans_11.trans_block.norm2.bias', 'conv_trans_11.trans_block.mlp.fc1.weight', 'conv_trans_11.trans_block.mlp.fc1.bias', 'conv_trans_11.trans_block.mlp.fc2.weight', 'conv_trans_11.trans_block.mlp.fc2.bias', 'conv_trans_12.squeeze_block.conv_project.weight', 'conv_trans_12.squeeze_block.conv_project.bias', 'conv_trans_12.squeeze_block.ln.weight', 'conv_trans_12.squeeze_block.ln.bias', 'conv_trans_12.expand_block.conv_project.weight', 'conv_trans_12.trans_block.norm1.weight', 'conv_trans_12.trans_block.norm1.bias', 'conv_trans_12.trans_block.attn.qkv.weight', 'conv_trans_12.trans_block.attn.proj.weight', 'conv_trans_12.trans_block.attn.proj.bias', 'conv_trans_12.trans_block.norm2.weight', 'conv_trans_12.trans_block.norm2.bias', 'conv_trans_12.trans_block.mlp.fc1.weight', 'conv_trans_12.trans_block.mlp.fc1.bias', 'conv_trans_12.trans_block.mlp.fc2.weight', 'conv_trans_12.trans_block.mlp.fc2.bias'], unexpected_keys=[])"},"metadata":{}}],"execution_count":26},{"cell_type":"code","source":"model.load_state_dict(checkpoint,strict=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T16:16:48.976063Z","iopub.execute_input":"2024-11-18T16:16:48.976535Z","iopub.status.idle":"2024-11-18T16:16:49.075506Z","shell.execute_reply.started":"2024-11-18T16:16:48.976489Z","shell.execute_reply":"2024-11-18T16:16:49.072505Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mRuntimeError\u001b[0m                              Traceback (most recent call last)","Cell \u001b[0;32mIn[28], line 1\u001b[0m\n\u001b[0;32m----> 1\u001b[0m \u001b[43mmodel\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mload_state_dict\u001b[49m\u001b[43m(\u001b[49m\u001b[43mcheckpoint\u001b[49m\u001b[43m,\u001b[49m\u001b[43mstrict\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[38;5;28;43;01mFalse\u001b[39;49;00m\u001b[43m)\u001b[49m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/nn/modules/module.py:2215\u001b[0m, in \u001b[0;36mModule.load_state_dict\u001b[0;34m(self, state_dict, strict, assign)\u001b[0m\n\u001b[1;32m   2210\u001b[0m         error_msgs\u001b[38;5;241m.\u001b[39minsert(\n\u001b[1;32m   2211\u001b[0m             \u001b[38;5;241m0\u001b[39m, \u001b[38;5;124m'\u001b[39m\u001b[38;5;124mMissing key(s) in state_dict: \u001b[39m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m. \u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;241m.\u001b[39mformat(\n\u001b[1;32m   2212\u001b[0m                 \u001b[38;5;124m'\u001b[39m\u001b[38;5;124m, \u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;241m.\u001b[39mjoin(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mk\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124m'\u001b[39m \u001b[38;5;28;01mfor\u001b[39;00m k \u001b[38;5;129;01min\u001b[39;00m missing_keys)))\n\u001b[1;32m   2214\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mlen\u001b[39m(error_msgs) \u001b[38;5;241m>\u001b[39m \u001b[38;5;241m0\u001b[39m:\n\u001b[0;32m-> 2215\u001b[0m     \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mRuntimeError\u001b[39;00m(\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mError(s) in loading state_dict for \u001b[39m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m:\u001b[39m\u001b[38;5;130;01m\\n\u001b[39;00m\u001b[38;5;130;01m\\t\u001b[39;00m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m'\u001b[39m\u001b[38;5;241m.\u001b[39mformat(\n\u001b[1;32m   2216\u001b[0m                        \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m\u001b[38;5;18m__class__\u001b[39m\u001b[38;5;241m.\u001b[39m\u001b[38;5;18m__name__\u001b[39m, \u001b[38;5;124m\"\u001b[39m\u001b[38;5;130;01m\\n\u001b[39;00m\u001b[38;5;130;01m\\t\u001b[39;00m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;241m.\u001b[39mjoin(error_msgs)))\n\u001b[1;32m   2217\u001b[0m \u001b[38;5;28;01mreturn\u001b[39;00m _IncompatibleKeys(missing_keys, unexpected_keys)\n","\u001b[0;31mRuntimeError\u001b[0m: Error(s) in loading state_dict for Conformer:\n\tsize mismatch for cls_token: copying a param with shape torch.Size([1, 1, 384]) from checkpoint, the shape in current model is torch.Size([1, 1, 768]).\n\tsize mismatch for trans_norm.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_norm.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_cls_head.weight: copying a param with shape torch.Size([1000, 384]) from checkpoint, the shape in current model is torch.Size([1000, 768]).\n\tsize mismatch for trans_patch_conv.weight: copying a param with shape torch.Size([384, 64, 4, 4]) from checkpoint, the shape in current model is torch.Size([768, 64, 4, 4]).\n\tsize mismatch for trans_patch_conv.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for trans_1.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for trans_1.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for trans_1.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for trans_1.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for trans_1.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 64, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 64, 1, 1]).\n\tsize mismatch for conv_trans_2.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.expand_block.conv_project.weight: copying a param with shape torch.Size([64, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([64, 768, 1, 1]).\n\tsize mismatch for conv_trans_2.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_2.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_2.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_2.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_2.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_2.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 64, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 64, 1, 1]).\n\tsize mismatch for conv_trans_3.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.expand_block.conv_project.weight: copying a param with shape torch.Size([64, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([64, 768, 1, 1]).\n\tsize mismatch for conv_trans_3.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_3.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_3.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_3.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_3.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_3.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 64, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 64, 1, 1]).\n\tsize mismatch for conv_trans_4.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.expand_block.conv_project.weight: copying a param with shape torch.Size([64, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([64, 768, 1, 1]).\n\tsize mismatch for conv_trans_4.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_4.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_4.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_4.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_4.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_4.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 128, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 128, 1, 1]).\n\tsize mismatch for conv_trans_5.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.expand_block.conv_project.weight: copying a param with shape torch.Size([128, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([128, 768, 1, 1]).\n\tsize mismatch for conv_trans_5.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_5.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_5.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_5.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_5.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_5.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 128, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 128, 1, 1]).\n\tsize mismatch for conv_trans_6.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.expand_block.conv_project.weight: copying a param with shape torch.Size([128, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([128, 768, 1, 1]).\n\tsize mismatch for conv_trans_6.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_6.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_6.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_6.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_6.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_6.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 128, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 128, 1, 1]).\n\tsize mismatch for conv_trans_7.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.expand_block.conv_project.weight: copying a param with shape torch.Size([128, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([128, 768, 1, 1]).\n\tsize mismatch for conv_trans_7.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_7.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_7.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_7.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_7.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_7.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 128, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 128, 1, 1]).\n\tsize mismatch for conv_trans_8.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.expand_block.conv_project.weight: copying a param with shape torch.Size([128, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([128, 768, 1, 1]).\n\tsize mismatch for conv_trans_8.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_8.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_8.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_8.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_8.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_8.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 256, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 256, 1, 1]).\n\tsize mismatch for conv_trans_9.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.expand_block.conv_project.weight: copying a param with shape torch.Size([256, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([256, 768, 1, 1]).\n\tsize mismatch for conv_trans_9.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_9.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_9.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_9.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_9.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_9.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 256, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 256, 1, 1]).\n\tsize mismatch for conv_trans_10.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.expand_block.conv_project.weight: copying a param with shape torch.Size([256, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([256, 768, 1, 1]).\n\tsize mismatch for conv_trans_10.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_10.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_10.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_10.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_10.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_10.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 256, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 256, 1, 1]).\n\tsize mismatch for conv_trans_11.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.expand_block.conv_project.weight: copying a param with shape torch.Size([256, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([256, 768, 1, 1]).\n\tsize mismatch for conv_trans_11.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_11.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_11.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_11.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_11.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_11.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 256, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 256, 1, 1]).\n\tsize mismatch for conv_trans_12.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.expand_block.conv_project.weight: copying a param with shape torch.Size([256, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([256, 768, 1, 1]).\n\tsize mismatch for conv_trans_12.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_12.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_12.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_12.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_12.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_12.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768])."],"ename":"RuntimeError","evalue":"Error(s) in loading state_dict for Conformer:\n\tsize mismatch for cls_token: copying a param with shape torch.Size([1, 1, 384]) from checkpoint, the shape in current model is torch.Size([1, 1, 768]).\n\tsize mismatch for trans_norm.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_norm.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_cls_head.weight: copying a param with shape torch.Size([1000, 384]) from checkpoint, the shape in current model is torch.Size([1000, 768]).\n\tsize mismatch for trans_patch_conv.weight: copying a param with shape torch.Size([384, 64, 4, 4]) from checkpoint, the shape in current model is torch.Size([768, 64, 4, 4]).\n\tsize mismatch for trans_patch_conv.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for trans_1.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for trans_1.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for trans_1.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for trans_1.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for trans_1.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for trans_1.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 64, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 64, 1, 1]).\n\tsize mismatch for conv_trans_2.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.expand_block.conv_project.weight: copying a param with shape torch.Size([64, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([64, 768, 1, 1]).\n\tsize mismatch for conv_trans_2.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_2.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_2.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_2.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_2.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_2.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_2.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 64, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 64, 1, 1]).\n\tsize mismatch for conv_trans_3.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.expand_block.conv_project.weight: copying a param with shape torch.Size([64, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([64, 768, 1, 1]).\n\tsize mismatch for conv_trans_3.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_3.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_3.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_3.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_3.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_3.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_3.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 64, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 64, 1, 1]).\n\tsize mismatch for conv_trans_4.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.expand_block.conv_project.weight: copying a param with shape torch.Size([64, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([64, 768, 1, 1]).\n\tsize mismatch for conv_trans_4.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_4.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_4.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_4.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_4.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_4.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_4.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 128, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 128, 1, 1]).\n\tsize mismatch for conv_trans_5.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.expand_block.conv_project.weight: copying a param with shape torch.Size([128, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([128, 768, 1, 1]).\n\tsize mismatch for conv_trans_5.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_5.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_5.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_5.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_5.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_5.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_5.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 128, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 128, 1, 1]).\n\tsize mismatch for conv_trans_6.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.expand_block.conv_project.weight: copying a param with shape torch.Size([128, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([128, 768, 1, 1]).\n\tsize mismatch for conv_trans_6.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_6.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_6.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_6.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_6.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_6.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_6.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 128, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 128, 1, 1]).\n\tsize mismatch for conv_trans_7.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.expand_block.conv_project.weight: copying a param with shape torch.Size([128, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([128, 768, 1, 1]).\n\tsize mismatch for conv_trans_7.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_7.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_7.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_7.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_7.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_7.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_7.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 128, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 128, 1, 1]).\n\tsize mismatch for conv_trans_8.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.expand_block.conv_project.weight: copying a param with shape torch.Size([128, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([128, 768, 1, 1]).\n\tsize mismatch for conv_trans_8.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_8.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_8.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_8.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_8.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_8.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_8.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 256, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 256, 1, 1]).\n\tsize mismatch for conv_trans_9.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.expand_block.conv_project.weight: copying a param with shape torch.Size([256, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([256, 768, 1, 1]).\n\tsize mismatch for conv_trans_9.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_9.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_9.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_9.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_9.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_9.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_9.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 256, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 256, 1, 1]).\n\tsize mismatch for conv_trans_10.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.expand_block.conv_project.weight: copying a param with shape torch.Size([256, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([256, 768, 1, 1]).\n\tsize mismatch for conv_trans_10.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_10.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_10.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_10.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_10.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_10.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_10.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 256, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 256, 1, 1]).\n\tsize mismatch for conv_trans_11.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.expand_block.conv_project.weight: copying a param with shape torch.Size([256, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([256, 768, 1, 1]).\n\tsize mismatch for conv_trans_11.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_11.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_11.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_11.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_11.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_11.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_11.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.squeeze_block.conv_project.weight: copying a param with shape torch.Size([384, 256, 1, 1]) from checkpoint, the shape in current model is torch.Size([768, 256, 1, 1]).\n\tsize mismatch for conv_trans_12.squeeze_block.conv_project.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.squeeze_block.ln.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.squeeze_block.ln.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.expand_block.conv_project.weight: copying a param with shape torch.Size([256, 384, 1, 1]) from checkpoint, the shape in current model is torch.Size([256, 768, 1, 1]).\n\tsize mismatch for conv_trans_12.trans_block.norm1.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.norm1.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.attn.qkv.weight: copying a param with shape torch.Size([1152, 384]) from checkpoint, the shape in current model is torch.Size([2304, 768]).\n\tsize mismatch for conv_trans_12.trans_block.attn.proj.weight: copying a param with shape torch.Size([384, 384]) from checkpoint, the shape in current model is torch.Size([768, 768]).\n\tsize mismatch for conv_trans_12.trans_block.attn.proj.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.norm2.weight: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.norm2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).\n\tsize mismatch for conv_trans_12.trans_block.mlp.fc1.weight: copying a param with shape torch.Size([1536, 384]) from checkpoint, the shape in current model is torch.Size([3072, 768]).\n\tsize mismatch for conv_trans_12.trans_block.mlp.fc1.bias: copying a param with shape torch.Size([1536]) from checkpoint, the shape in current model is torch.Size([3072]).\n\tsize mismatch for conv_trans_12.trans_block.mlp.fc2.weight: copying a param with shape torch.Size([384, 1536]) from checkpoint, the shape in current model is torch.Size([768, 3072]).\n\tsize mismatch for conv_trans_12.trans_block.mlp.fc2.bias: copying a param with shape torch.Size([384]) from checkpoint, the shape in current model is torch.Size([768]).","output_type":"error"}],"execution_count":28},{"cell_type":"code","source":"# Load the checkpoint (assuming 'checkpoint' is your pretrained model's checkpoint dictionary)\ncheckpoint = torch.load(\"path/to/your/checkpoint.pth\")\n\n# Get the current state dict of your modified model\nmodel_state_dict = model.state_dict()\n\n# Filter out incompatible keys\nfiltered_checkpoint = {k: v for k, v in checkpoint.items() if k in model_state_dict and model_state_dict[k].shape == v.shape}\n\n# Load only the compatible layers\nmodel.load_state_dict(filtered_checkpoint, strict=False)\n","metadata":{"trusted":true},"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.66748Z","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}]}