{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":30665,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nfrom glob import glob\nfrom tqdm.notebook import tqdm\nimport joblib\nimport time\n\nimport matplotlib.pyplot as plt \n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport torchaudio\nimport torchvision\nfrom torch.utils.data import Dataset, DataLoader\nfrom typing import Any, Callable, List, Optional, Type, Union\nfrom torch import Tensor","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-25T15:44:18.022016Z","iopub.execute_input":"2024-03-25T15:44:18.022501Z","iopub.status.idle":"2024-03-25T15:44:18.03085Z","shell.execute_reply.started":"2024-03-25T15:44:18.022465Z","shell.execute_reply":"2024-03-25T15:44:18.029394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading data\n## data path","metadata":{}},{"cell_type":"code","source":"BASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\n\nSPEC_DIR = \"/tmp/dataset/hms-hbac\"\nos.makedirs(SPEC_DIR+'/train_spectrograms', exist_ok=True)\nos.makedirs(SPEC_DIR+'/test_spectrograms', exist_ok=True)\n\nclass_names = ['Seizure', 'LPD', 'GPD', 'LRDA','GRDA', 'Other']\nlabel2name = dict(enumerate(class_names))\nname2label = {v:k for k, v in label2name.items()}","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:33:23.390054Z","iopub.execute_input":"2024-03-25T15:33:23.39113Z","iopub.status.idle":"2024-03-25T15:33:23.408002Z","shell.execute_reply.started":"2024-03-25T15:33:23.391095Z","shell.execute_reply":"2024-03-25T15:33:23.406694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train + Valid\nsample_rate=200\ndf = pd.read_csv(f'{BASE_PATH}/train.csv')\ndf['eeg_path'] = f'{BASE_PATH}/train_eegs/'+df['eeg_id'].astype(str)+'.parquet'\ndf['spec_path'] = f'{BASE_PATH}/train_spectrograms/'+df['spectrogram_id'].astype(str)+'.parquet'\n# df['spec_sub_path'] = f'{BASE_PATH}/train_spectrograms/'+df['spectrogram_sub_id'].astype(str)+'.parquet'\ndf['spec2_path'] = f'{SPEC_DIR}/train_spectrograms/'+df['spectrogram_id'].astype(str)+'.npy'\ndf['class_name'] = df.expert_consensus.copy()\ndf['class_label'] = df.expert_consensus.map(name2label)\ndisplay(df.head(2))\n\n# Test\ntest_df = pd.read_csv(f'{BASE_PATH}/test.csv')\ntest_df['eeg_path'] = f'{BASE_PATH}/test_eegs/'+test_df['eeg_id'].astype(str)+'.parquet'\ntest_df['spec_path'] = f'{BASE_PATH}/test_spectrograms/'+test_df['spectrogram_id'].astype(str)+'.parquet'\ntest_df['spec2_path'] = f'{SPEC_DIR}/test_spectrograms/'+test_df['spectrogram_id'].astype(str)+'.npy'\ndisplay(test_df.head(2))","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:33:23.409986Z","iopub.execute_input":"2024-03-25T15:33:23.410501Z","iopub.status.idle":"2024-03-25T15:33:24.122338Z","shell.execute_reply.started":"2024-03-25T15:33:23.410458Z","shell.execute_reply":"2024-03-25T15:33:24.121155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx = 30\nprint(df['eeg_id'][idx])\neeg_path = df['eeg_path'][idx]\neeg = pd.read_parquet(eeg_path)\nprint(eeg.tail(5))\n\nprint(df['spectrogram_id'][idx])\nspec_path = df['spec_path'][idx]\nspec = pd.read_parquet(spec_path)\nprint(spec.tail(5))","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:33:24.125889Z","iopub.execute_input":"2024-03-25T15:33:24.126392Z","iopub.status.idle":"2024-03-25T15:33:24.39185Z","shell.execute_reply.started":"2024-03-25T15:33:24.126347Z","shell.execute_reply":"2024-03-25T15:33:24.390375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def cut_egg(egg, offest):\n    offset *= sample_rate\n    length = 50*sample_rate\n    return egg[offset: offset+length]\n\n# n_frame = (N-w)/step +1\n#         = (600 + L - 50)/2\n#        = L/2 + 275torch.tensor(spec.values)\ndef cut_spec(spec, offset):\n    offset = int(offset/2)\n    length = 275+ int((offset+50)/2)\n    return spec[offset: offset+length]","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:33:24.393323Z","iopub.execute_input":"2024-03-25T15:33:24.393754Z","iopub.status.idle":"2024-03-25T15:33:24.401264Z","shell.execute_reply.started":"2024-03-25T15:33:24.393715Z","shell.execute_reply":"2024-03-25T15:33:24.399762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spec_corpus={}\ndef get_spec(spec_id, split=\"train\",):\n    spec_path = f'{BASE_PATH}/{split}_spectrograms/'+spec_id.astype(str)+'.parquet'\n    spec = pd.read_parquet(spec_path)\n    spec = spec.fillna(0)  # fill NaN values with 0, \n    spec = torch.tensor(spec.values, dtype=torch.float64)[:,1:]\n    torch.save(spec, f\"{SPEC_DIR}/{split}_spectrograms/{spec_id}.pt\")\n    return {spec_id: spec}\n\n# # Get unique spec_ids of train and valid data\nspec_ids = df[\"spectrogram_id\"].unique()\n\n# Parallelize the processing using joblib for training data\ntrain_data = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(get_spec)(spec_id, \"train\")\n    for spec_id in tqdm(spec_ids, total=len(spec_ids))\n)\nfor spec in train_data:\n    spec_corpus.update(spec)\n\n# Get unique spec_ids of test data\ntest_spec_ids = test_df[\"spectrogram_id\"].unique()\n\n# Parallelize the processing using joblib for test data\ntest_data = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(get_spec)(spec_id, \"test\")\n    for spec_id in tqdm(test_spec_ids, total=len(test_spec_ids))\n)\nfor spec in test_data:\n    spec_corpus.update(spec)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:33:24.402789Z","iopub.execute_input":"2024-03-25T15:33:24.403153Z","iopub.status.idle":"2024-03-25T15:38:08.14128Z","shell.execute_reply.started":"2024-03-25T15:33:24.403124Z","shell.execute_reply":"2024-03-25T15:38:08.139385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pytorch Dataset and Dataloader","metadata":{}},{"cell_type":"code","source":"class HBACDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        spec_id = self.df['spectrogram_id'][idx]\n        spec = spec_corpus[spec_id]\n        if 'spectrogram_label_offset_seconds' in self.df:\n            spec = cut_spec(spec, self.df['spectrogram_label_offset_seconds'][idx] )\n        \n        # seizure_vote\tlpd_vote\tgpd_vote\tlrda_vote\tgrda_vote\tother_vote\n        try:\n            # the more numbervote, the true prob more approximately 1.0\n            target = torch.tensor([(df[label+\"_vote\"][idx] +0.05) for label in [\"seizure\", \"lpd\", \"gpd\", \"lrda\", \"grda\", \"other\"]])\n            target = target / target.sum()  # probality \n        except:\n            # test mode\n            target = torch.zeros(6,) \n        \n        if self.transform:\n            spec = self.transform(spec)\n\n        return spec, target\n    \nclass Transform(nn.Module):\n    def __init__(self, prob_freq_masking=0.2, freq_mask_param =50, prob_time_masking=0.2, time_mask_param=50):\n        super().__init__()\n        self.prob_freq_masking = prob_freq_masking\n        self.freq_masking = torchaudio.transforms.FrequencyMasking(freq_mask_param=freq_mask_param)\n        \n        self.prob_time_masking = prob_time_masking\n        self.time_masking = torchaudio.transforms.TimeMasking(time_mask_param=time_mask_param)\n        \n    def forward(self, spectrogram):\n        prob = torch.rand(2)\n        if prob[0] < self.prob_freq_masking:\n            spectrogram = self.freq_masking(spectrogram)\n        if prob[0] < self.prob_time_masking:\n            spectrogram = self.time_masking(spectrogram)\n        return spectrogram\n\ndef collate_spec(batch):\n    spec, target = zip(*batch)\n    spec_length = torch.tensor([len(sample) for sample in spec])\n    spec = nn.utils.rnn.pad_sequence([sample for sample in spec], batch_first=True, padding_value=0.0)\n    target = torch.stack([t for t in target], dim =0)\n    \n    return spec, spec_length, target\n    ","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:38:08.144257Z","iopub.execute_input":"2024-03-25T15:38:08.145078Z","iopub.status.idle":"2024-03-25T15:38:08.164562Z","shell.execute_reply.started":"2024-03-25T15:38:08.145035Z","shell.execute_reply":"2024-03-25T15:38:08.163139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Split Train and Valid","metadata":{}},{"cell_type":"code","source":"spec_ids = df[\"spectrogram_id\"].unique()\nnp.random.shuffle(spec_ids)\n\n# Sample from full data\nsample_df = df.groupby(\"spectrogram_id\").head(1).reset_index(drop=True)\n\ntrain_size= int(len(sample_df)*0.8)\ndf_train = sample_df.sample(train_size)\ndf_val = sample_df.drop(df_train.index).reset_index(drop=True)\ndf_train = df_train.reset_index(drop=True)\n\nprint(len(df_train))\nprint(len(df_val))","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:38:08.16616Z","iopub.execute_input":"2024-03-25T15:38:08.166546Z","iopub.status.idle":"2024-03-25T15:38:08.270475Z","shell.execute_reply.started":"2024-03-25T15:38:08.166515Z","shell.execute_reply":"2024-03-25T15:38:08.269222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# transform = Transform()\n\nbatch_size=32\nnum_workers=4\n\ntransform =  None\n#Transform()\n\ntrain_ds = HBACDataset(df_train, transform)\ntrain_dl = DataLoader(train_ds, batch_size=batch_size, num_workers=num_workers ,shuffle=True,collate_fn=collate_spec)\n\nval_ds = HBACDataset(df_val, transform)\nval_dl = DataLoader(val_ds, batch_size=batch_size, num_workers=num_workers ,shuffle=True,collate_fn=collate_spec)\n\ntest_ds = HBACDataset(test_df)\ntest_dl = DataLoader(test_ds, batch_size=batch_size, num_workers=num_workers, shuffle=False,collate_fn=collate_spec)","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:38:08.272405Z","iopub.execute_input":"2024-03-25T15:38:08.273335Z","iopub.status.idle":"2024-03-25T15:38:08.28278Z","shell.execute_reply.started":"2024-03-25T15:38:08.273282Z","shell.execute_reply":"2024-03-25T15:38:08.281169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":" \nclass AdaptCNN(nn.Module):\n    def __init__(self, \n                 input_channels=4,\n                 c_out_1 = 16, \n                 c_out_2 = 32,\n                 c_out_3 = 64,\n                 kernel_size=(3,3), \n                 dropout=0.2,\n                 pool_1=[24,7],\n                 pool_2=[12,5],\n                 pool_3=[6,3]\n                 ):\n        super().__init__()\n        self.input_channels = input_channels\n        self.c_out_1 = c_out_1\n        self.c_out_2 = c_out_2\n        self.c_out_3 = c_out_3\n        self.kernel_size = kernel_size\n        self.pool_1 = pool_1\n        self.pool_2 = pool_2\n        self.pool_3 = pool_3\n        self.dropout_rate = dropout\n\n        self.dropout = nn.Dropout2d(p=self.dropout_rate)\n        \n        if isinstance(self.kernel_size, int):\n            self.kernel_size = (self.kernel_size, self.kernel_size)\n            \n        # Set kernel width of last conv layer to last pool width to \n        # downsample width to one.\n        self.kernel_size_last = (self.kernel_size[0], self.pool_3[1])\n            \n        # kernel_size[1]=1 can be used for seg_length=1 -> corresponds to \n        # 1D conv layer, no width padding needed.\n        if self.kernel_size[1] == 1:\n            self.cnn_pad = (1,0)\n        else:\n            self.cnn_pad = (1,1)   \n            \n        self.conv1 = nn.Conv2d(self.input_channels, self.c_out_1, self.kernel_size, padding = self.cnn_pad)\n        self.bn1 = nn.BatchNorm2d( self.conv1.out_channels )\n\n        self.conv2 = nn.Conv2d(self.conv1.out_channels, self.c_out_2, self.kernel_size, padding = self.cnn_pad)\n        self.bn2 = nn.BatchNorm2d( self.conv2.out_channels )\n\n        self.conv3 = nn.Conv2d(self.conv2.out_channels, self.c_out_3, self.kernel_size, padding = self.cnn_pad)\n        self.bn3 = nn.BatchNorm2d( self.conv3.out_channels )\n\n        self.conv4 = nn.Conv2d(self.conv3.out_channels, self.c_out_3, self.kernel_size, padding = self.cnn_pad)\n        self.bn4 = nn.BatchNorm2d( self.conv4.out_channels )\n\n        self.conv5 = nn.Conv2d(self.conv4.out_channels, self.c_out_3, self.kernel_size, padding = self.cnn_pad)\n        self.bn5 = nn.BatchNorm2d( self.conv5.out_channels )\n\n        self.conv6 = nn.Conv2d(self.conv5.out_channels, self.c_out_3, self.kernel_size_last, padding = (1,0))\n        self.bn6 = nn.BatchNorm2d( self.conv6.out_channels )\n        \n        self.fan_out = (self.conv6.out_channels * self.pool_3[0])\n\n    def forward(self, x):\n        x = F.relu( self.bn1( self.conv1(x) ) )\n        x = F.adaptive_max_pool2d(x, output_size=(self.pool_1))\n\n        x = F.relu( self.bn2( self.conv2(x) ) )\n        x = F.adaptive_max_pool2d(x, output_size=(self.pool_2))\n        \n        x = self.dropout(x)\n        x = F.relu( self.bn3( self.conv3(x) ) )\n        x = self.dropout(x)\n        x = F.relu( self.bn4( self.conv4(x) ) )\n        x = F.adaptive_max_pool2d(x, output_size=(self.pool_3))\n\n        x = self.dropout(x)\n        x = F.relu( self.bn5( self.conv5(x) ) )\n        x = self.dropout(x)\n        x = F.relu( self.bn6( self.conv6(x) ) )\n        x = x.view(-1, self.conv6.out_channels * self.pool_3[0])\n\n        return x\n\n\n\nclass SelfAttention(nn.Module):\n    '''\n    SelfAttention: The main SelfAttention module that can be used as a\n    time-dependency model.                                            \n    '''         \n    def __init__(self,\n                 input_size,        # cnn.model.fan_out\n                 d_model=512,\n                 nhead=8,\n                 pool_size=3,\n                 pos_enc=None,\n                 num_layers=6,\n                 sa_h=2048,\n                 dropout=0.1,\n                 ):\n        super().__init__()\n        \n        encoder_layer = SelfAttentionLayer(d_model, nhead, pool_size, sa_h, dropout)\n        self.norm1 = nn.LayerNorm(d_model)\n        \n        self.linear = nn.Linear(input_size, d_model)\n        \n        self.layers = self._get_clones(encoder_layer, num_layers)\n        self.num_layers = num_layers\n        self.d_model = d_model\n        self.nhead = nhead      \n        \n        if pos_enc:\n            self.pos_encoder = PositionalEncoding(d_model, dropout)\n        else:\n            self.pos_encoder = nn.Identity()\n            \n        self._reset_parameters()\n        \n    def _get_clones(self, module, N):\n        return nn.ModuleList([copy.deepcopy(module) for i in range(N)])\n    \n    def _reset_parameters(self):\n        for p in self.parameters():\n            if p.dim() > 1:\n                nn.init.xavier_uniform_(p)\n\n    def forward(self, src, n_wins=None):            \n        src = self.linear(src)\n        output = src.transpose(1,0)\n        output = self.norm1(output)\n        output = self.pos_encoder(output)\n        \n        for mod in self.layers:\n            output, n_wins = mod(output, n_wins=n_wins)\n        return output.transpose(1,0), n_wins\n    \nclass SelfAttentionLayer(nn.Module):\n    '''\n    SelfAttentionLayer: The SelfAttentionLayer that is used by the\n    SelfAttention module.                                            \n    '''          \n    def __init__(self, d_model, nhead, pool_size=1, sa_h=2048, dropout=0.1):\n        super().__init__()\n        \n        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)\n        \n        self.linear1 = nn.Linear(d_model, sa_h)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(sa_h, d_model)\n\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n        \n        self.activation = F.relu          \n        \n    def forward(self, src, n_wins=None):\n        \n        if n_wins is not None:\n            mask = ~((torch.arange(src.shape[0])[None, :]).to(src.device) < n_wins[:, None].to(torch.long).to(src.device))\n        else:\n            mask = None\n        \n        src2 = self.self_attn(src, src, src, key_padding_mask=mask)[0]\n        src = src + self.dropout1(src2)\n        src = self.norm1(src)\n        src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))\n        src = src + self.dropout2(src2)\n\n        src = self.norm2(src)\n\n        return src, n_wins\n    \nclass PoolAttFF(nn.Module):\n    '''\n    PoolAttFF: Attention-Pooling module with additonal feed-forward network.\n    '''         \n    def __init__(self, d_input, output_size, h, dropout=0.1):\n        super().__init__()\n        \n        self.linear1 = nn.Linear(d_input, h)\n        self.linear2 = nn.Linear(h, 1)\n        \n        self.linear3 = nn.Linear(d_input, output_size)\n        \n        self.activation = F.relu\n        self.dropout = nn.Dropout(dropout)\n        \n    def forward(self, x, n_wins):\n\n        att = self.linear2(self.dropout(self.activation(self.linear1(x))))\n        att = att.transpose(2,1)\n        mask = torch.arange(att.shape[2])[None, :] < n_wins[:, None].to('cpu').to(torch.long)\n        att[~mask.unsqueeze(1)] = float(\"-Inf\")          \n        att = F.softmax(att, dim=2)\n        x = torch.bmm(att, x) \n        x = x.squeeze(1)\n        \n        x = self.linear3(x)\n        \n        return x    ","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:38:08.287337Z","iopub.execute_input":"2024-03-25T15:38:08.287729Z","iopub.status.idle":"2024-03-25T15:38:08.334948Z","shell.execute_reply.started":"2024-03-25T15:38:08.287699Z","shell.execute_reply":"2024-03-25T15:38:08.333621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def conv3x3(in_planes: int, out_planes: int, stride: int = 1, groups: int = 1, dilation: int = 1) -> nn.Conv2d:\n    \"\"\"3x3 convolution with padding\"\"\"\n    return nn.Conv2d(\n        in_planes,\n        out_planes,\n        kernel_size=3,\n        stride=stride,\n        padding=dilation,\n        groups=groups,\n        bias=False,\n        dilation=dilation,\n    )\n\ndef conv1x1(in_planes: int, out_planes: int, stride: int = 1) -> nn.Conv2d:\n    \"\"\"1x1 convolution\"\"\"\n    return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)\n\nclass BasicBlock(nn.Module):\n    expansion: int = 1\n    def __init__(\n        self,\n        inplanes: int,\n        planes: int,\n        stride: int = 1,\n        downsample: Optional[nn.Module] = None,\n        groups: int = 1,\n        base_width: int = 64,\n        dilation: int = 1,\n        norm_layer: Optional[Callable[..., nn.Module]] = None,\n    ) -> None:\n        super().__init__()\n        if norm_layer is None:\n            norm_layer = nn.BatchNorm2d\n        if groups != 1 or base_width != 64:\n            raise ValueError(\"BasicBlock only supports groups=1 and base_width=64\")\n        if dilation > 1:\n            raise NotImplementedError(\"Dilation > 1 not supported in BasicBlock\")\n        # Both self.conv1 and self.downsample layers downsample the input when stride != 1\n        self.conv1 = conv3x3(inplanes, planes, stride)\n        self.bn1 = norm_layer(planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv3x3(planes, planes)\n        self.bn2 = norm_layer(planes)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x: Tensor) -> Tensor:\n        identity = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        if self.downsample is not None:\n            identity = self.downsample(x)\n\n        out += identity\n        out = self.relu(out)\n\n        return out\n\n\nclass Bottleneck(nn.Module):\n    # Bottleneck in torchvision places the stride for downsampling at 3x3 convolution(self.conv2)\n    # while original implementation places the stride at the first 1x1 convolution(self.conv1)\n    # according to \"Deep residual learning for image recognition\" https://arxiv.org/abs/1512.03385.\n    # This variant is also known as ResNet V1.5 and improves accuracy according to\n    # https://ngc.nvidia.com/catalog/model-scripts/nvidia:resnet_50_v1_5_for_pytorch.\n    expansion: int = 4\n    def __init__(\n        self,\n        inplanes: int,\n        planes: int,\n        stride: int = 1,\n        downsample: Optional[nn.Module] = None,\n        groups: int = 1,\n        base_width: int = 64,\n        dilation: int = 1,\n        norm_layer: Optional[Callable[..., nn.Module]] = None,\n    ) -> None:\n        super().__init__()\n        if norm_layer is None:\n            norm_layer = nn.BatchNorm2d\n        width = int(planes * (base_width / 64.0)) * groups\n        # Both self.conv2 and self.downsample layers downsample the input when stride != 1\n        self.conv1 = conv1x1(inplanes, width)\n        self.bn1 = norm_layer(width)\n        self.conv2 = conv3x3(width, width, stride, groups, dilation)\n        self.bn2 = norm_layer(width)\n        self.conv3 = conv1x1(width, planes * self.expansion)\n        self.bn3 = norm_layer(planes * self.expansion)\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x: Tensor) -> Tensor:\n        identity = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n\n        out = self.conv3(out)\n        out = self.bn3(out)\n\n        if self.downsample is not None:\n            identity = self.downsample(x)\n\n        out += identity\n        out = self.relu(out)\n\n        return out\n\nclass ResNet(nn.Module):\n    def __init__(\n        self,\n        block: Type[Union[BasicBlock, Bottleneck]],\n        layers: List[int],\n        num_classes: int = 6,\n        zero_init_residual: bool = False,\n        groups: int = 1,\n        width_per_group: int = 64,\n        replace_stride_with_dilation: Optional[List[bool]] = None,\n        norm_layer: Optional[Callable[..., nn.Module]] = None,\n    ) -> None:\n        super().__init__()\n        if norm_layer is None:\n            norm_layer = nn.BatchNorm2d\n        self._norm_layer = norm_layer\n\n        self.inplanes = 64\n        self.dilation = 1\n        if replace_stride_with_dilation is None:\n            # each element in the tuple indicates if we should replace\n            # the 2x2 stride with a dilated convolution instead\n            replace_stride_with_dilation = [False, False, False]\n        if len(replace_stride_with_dilation) != 3:\n            raise ValueError(\n                \"replace_stride_with_dilation should be None \"\n                f\"or a 3-element tuple, got {replace_stride_with_dilation}\"\n            )\n        self.groups = groups\n        self.base_width = width_per_group\n        self.conv1 = nn.Conv2d(4, self.inplanes, kernel_size=7, stride=2, padding=3, bias=False)\n        self.bn1 = self._norm_layer(self.inplanes)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        self.layer1 = self._make_layer(block, 64, layers[0])\n        self.layer2 = self._make_layer(block, 128, layers[1], stride=2, dilate=replace_stride_with_dilation[0])\n        self.layer3 = self._make_layer(block, 256, layers[2], stride=2, dilate=replace_stride_with_dilation[1])\n        self.layer4 = self._make_layer(block, 512, layers[3], stride=2, dilate=replace_stride_with_dilation[2])\n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n        self.fc = nn.Linear(512 * block.expansion, num_classes)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode=\"fan_out\", nonlinearity=\"relu\")\n            elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n\n        # Zero-initialize the last BN in each residual branch,\n        # so that the residual branch starts with zeros, and each residual block behaves like an identity.\n        # This improves the model by 0.2~0.3% according to https://arxiv.org/abs/1706.02677\n        if zero_init_residual:\n            for m in self.modules():\n                if isinstance(m, Bottleneck) and m.bn3.weight is not None:\n                    nn.init.constant_(m.bn3.weight, 0)  # type: ignore[arg-type]\n                elif isinstance(m, BasicBlock) and m.bn2.weight is not None:\n                    nn.init.constant_(m.bn2.weight, 0)  # type: ignore[arg-type]\n\n    def _make_layer(\n        self,\n        block: Type[Union[BasicBlock, Bottleneck]],\n        planes: int,\n        blocks: int,\n        stride: int = 1,\n        dilate: bool = False,\n    ) -> nn.Sequential:\n        norm_layer = self._norm_layer\n        downsample = None\n        previous_dilation = self.dilation\n        if dilate:\n            self.dilation *= stride\n            stride = 1\n        if stride != 1 or self.inplanes != planes * block.expansion:\n            downsample = nn.Sequential(\n                conv1x1(self.inplanes, planes * block.expansion, stride),\n                norm_layer(planes * block.expansion),\n            )\n\n        layers = []\n        layers.append(\n            block(self.inplanes, planes, stride, downsample, self.groups, self.base_width, previous_dilation, norm_layer)\n        )\n        self.inplanes = planes * block.expansion\n        for _ in range(1, blocks):\n            layers.append(\n                block(self.inplanes,\n                    planes,\n                    groups=self.groups,\n                    base_width=self.base_width,\n                    dilation=self.dilation,\n                    norm_layer=norm_layer,\n                )\n            )\n\n        return nn.Sequential(*layers)\n\n    def _forward_impl(self, x: Tensor) -> Tensor:\n        # See note [TorchScript super()]\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n\n        x = self.avgpool(x)\n        x = torch.flatten(x, 1)\n        x = self.fc(x)\n\n        return x\n\n    def forward(self, x: Tensor, x_len: Tensor = None) -> Tensor:\n        x = x.float()\n        # x: N x L x 400\n        x = x.transpose(1,2)\n        x = x.reshape(x.size(0), 4, 100, -1)\n        return self._forward_impl(x)\n\n\nclass Network(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        self.conv1 = nn.Conv2d(in_channels=4, out_channels=16, kernel_size=5, stride=1, padding=1)\n        self.bn1 = nn.BatchNorm2d(16)\n        self.conv2 = nn.Conv2d(in_channels=16, out_channels=32, kernel_size=5, stride=1, padding=1)\n        self.bn2 = nn.BatchNorm2d(32)\n        self.pool = nn.MaxPool2d(2,2)\n        self.conv4 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=5, stride=1, padding=1)\n        self.bn4 = nn.BatchNorm2d(64)\n        self.conv5 = nn.Conv2d(in_channels=64, out_channels=64, kernel_size=5, stride=1, padding=1)\n        self.bn5 = nn.BatchNorm2d(64)\n        self.conv6=nn.AdaptiveMaxPool2d(output_size=(20,20))        # 20 Hz and 10s\n        self.bn6 = nn.BatchNorm2d(64)\n        self.fc1 = nn.Linear(64*20*20, 6)\n\n    def forward(self, input, input_len=None):\n        input = input.float()\n        # x: N x L x 400\n        input = input.transpose(1,2)\n        input = input.reshape(input.size(0), 4, 100, -1)\n        \n        output = F.relu(self.bn1(self.conv1(input)))      \n        output = F.relu(self.bn2(self.conv2(output)))     \n        output = self.pool(output)                        \n        output = F.relu(self.bn4(self.conv4(output)))     \n        output = F.relu(self.bn5(self.conv5(output)))           \n        output = F.relu(self.bn6(self.conv6(output)))     \n        output = output.view(-1, 64*20*20)\n        output = self.fc1(output)\n        return output\n","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:57:13.495404Z","iopub.execute_input":"2024-03-25T15:57:13.49586Z","iopub.status.idle":"2024-03-25T15:57:13.558744Z","shell.execute_reply.started":"2024-03-25T15:57:13.495822Z","shell.execute_reply":"2024-03-25T15:57:13.557735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cpu')\n\nmodel = ResNet(Bottleneck, [3, 4, 6, 3])\nmodel.to(device)\nkl_loss = nn.KLDivLoss(reduction=\"batchmean\")\nce_loss = nn.CrossEntropyLoss(reduction='mean')\n\noptimizer = optim.AdamW(model.parameters(), lr=0.0001, betas=(0.9, 0.999), eps=1e-09, weight_decay=0.01)\nwarmup_steps= int(len(train_dl))     # 1 epoch\n# # scheduler = NoamLR(optimizer,warmup_steps=warmup_steps, d_model=512)\n# optimizer = optim.SGD(model.parameters(), lr= 0.001, weight_decay=0.01)\nepochs = 20","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:57:14.22663Z","iopub.execute_input":"2024-03-25T15:57:14.22739Z","iopub.status.idle":"2024-03-25T15:57:14.695825Z","shell.execute_reply.started":"2024-03-25T15:57:14.227354Z","shell.execute_reply":"2024-03-25T15:57:14.694425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    spec, spec_length, target = next(iter(train_dl))\n    # print(x_len)\n    # print(target)\n\n    print(f\"Feature batch shape: {spec.size()}\")\n    print(f\"Length batch shape: {spec_length.size()}\")\n    print(f\"Labels batch shape: {target.size()}\")\n\n    logit = model(spec.to(device), spec_length.to(device))\n    prob = F.log_softmax(logit,dim=1).clamp(min=-1e9,max=0)\n    print(torch.exp(prob))     \n    loss = kl_loss(prob, target.to(device))\n    print(loss)\n    print(ce_loss(logit, target))","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:57:20.348931Z","iopub.execute_input":"2024-03-25T15:57:20.349698Z","iopub.status.idle":"2024-03-25T15:57:25.39905Z","shell.execute_reply.started":"2024-03-25T15:57:20.349659Z","shell.execute_reply":"2024-03-25T15:57:25.397471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in tqdm(range(epochs)):\n    start_time = time.time()\n    batch_cnt = 0\n    train_loss = 0.0\n    model.train()\n    pbar = tqdm(iterable=batch_cnt, total=len(train_dl)+len(val_dl), ascii=\">—\",\n                    bar_format='{bar} {percentage:3.0f}%, {n_fmt}/{total_fmt}, {elapsed}<{remaining}{postfix}')\n    iters = len(train_dl)\n    for step, batch in enumerate(train_dl):\n        # Estimate batch ---------------------------------------------------\n        spec, spec_length, target = batch\n        spec = spec.to(device)\n        spec_length = spec_length.to(device)\n        target = target.to(device)\n\n        # Forward pass ----------------------------------------------------\n        logit = model(spec, spec_length)\n\n        # Loss ------------------------------------------------------------     \n        prob = F.log_softmax(logit,dim=1).clamp(min=-1e9,max=0)\n        loss_batch = kl_loss(prob, target)\n        # kl_loss(prob, target) + ce_loss(logit, target)\n\n        # Backprop  -------------------------------------------------------\n        loss_batch.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n\n        # Update total loss -----------------------------------------------\n        train_loss += loss_batch.item()\n        batch_cnt += 1\n        # Scheduler update    ---------------------------------------------\n        # heduler.step(epoch * iters + step)\n\n        pbar.set_postfix(loss=loss_batch.item())    # ,lr=scheduler.get_lr()\n        pbar.update()\n        \n    print(f'<---- Training loss: {train_loss/batch_cnt:.3f}---->')\n        \n    val_loss = 0.0\n    model.eval()\n    batch_cnt = 0\n    with torch.no_grad():\n        for batch in val_dl:\n            spec, spec_length, target = batch\n            spec = spec.to(device)\n            spec_length = spec_length.to(device)\n            target = target.to(device)\n\n            logit = model(spec, spec_length)\n            prob = F.log_softmax(logit,dim=1).clamp(min=-1e9,max=0)\n            loss_batch = kl_loss(prob, target) \n            \n            val_loss += loss_batch.item()\n            batch_cnt += 1\n\n            pbar.set_postfix(loss=loss_batch.item())\n            pbar.update()\n\n    print(f'<---- Validation loss: {val_loss/batch_cnt:.3f}---->')\n\n    elapsed_time = time.time() - start_time\n    print(\"The time elapse of epoch {:03d}\".format(epoch) + \" is: \" + \n            time.strftime(\"%H: %M: %S\", time.gmtime(elapsed_time)))","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:44:34.765468Z","iopub.execute_input":"2024-03-25T15:44:34.766189Z","iopub.status.idle":"2024-03-25T15:48:28.221667Z","shell.execute_reply.started":"2024-03-25T15:44:34.766144Z","shell.execute_reply":"2024-03-25T15:48:28.219617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = '/kaggle/working/conformer.ckpt'\ntorch.save(model.state_dict(), model_path)","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:38:11.084532Z","iopub.status.idle":"2024-03-25T15:38:11.085287Z","shell.execute_reply.started":"2024-03-25T15:38:11.085054Z","shell.execute_reply":"2024-03-25T15:38:11.085074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"preds = []\nmodel.eval()\nwith torch.no_grad():\n    for spec, spec_length, target in test_dl:\n        spec = spec.cuda()\n        spec_length = spec_length.cuda()\n        pred = model(spec, spec_length)\n        preds.append(pred)\npreds = torch.cat(preds, dim =0)","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:38:11.086928Z","iopub.status.idle":"2024-03-25T15:38:11.087683Z","shell.execute_reply.started":"2024-03-25T15:38:11.087452Z","shell.execute_reply":"2024-03-25T15:38:11.087472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_df = test_df[[\"eeg_id\"]].copy()\ntarget_cols = [x.lower()+'_vote' for x in class_names]\npred_df[target_cols] = preds.tolist()\n\nsub_df = pd.read_csv(f'{BASE_PATH}/sample_submission.csv')\nsub_df = sub_df[[\"eeg_id\"]].copy()\nsub_df = sub_df.merge(pred_df, on=\"eeg_id\", how=\"left\")\nsub_df.to_csv(\"submission.csv\", index=False)\nsub_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:38:11.089024Z","iopub.status.idle":"2024-03-25T15:38:11.089749Z","shell.execute_reply.started":"2024-03-25T15:38:11.08953Z","shell.execute_reply":"2024-03-25T15:38:11.089548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}