{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7898450,"sourceType":"datasetVersion","datasetId":4638387}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# imports","metadata":{}},{"cell_type":"code","source":"from datetime import datetime\nimport random\n# import einops\nimport wandb\nimport pandas as pd\nimport numpy as np\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt \nfrom torch.utils.data import Dataset, DataLoader, random_split, Subset\nimport torch\nimport torch as t\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom functools import lru_cache\n\ndevice = 'cuda' if t.cuda.is_available() else 'cpu'","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2024-03-29T18:57:28.868609Z","iopub.execute_input":"2024-03-29T18:57:28.868967Z","iopub.status.idle":"2024-03-29T18:57:34.421266Z","shell.execute_reply.started":"2024-03-29T18:57:28.868938Z","shell.execute_reply":"2024-03-29T18:57:34.420178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# utils","metadata":{}},{"cell_type":"code","source":"import gc \ndef GC():\n    gc.collect()\n    t.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-03-29T21:03:56.118030Z","iopub.execute_input":"2024-03-29T21:03:56.118424Z","iopub.status.idle":"2024-03-29T21:03:56.123348Z","shell.execute_reply.started":"2024-03-29T21:03:56.118389Z","shell.execute_reply":"2024-03-29T21:03:56.122133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# validation / inference loss\n@t.no_grad()\ndef eval(model, x, y, do_eval=True):\n    assert not x.isnan().any()\n    assert not y.isnan().any()\n    if do_eval: model.eval()\n    else: model.train()\n    logs = model(x.to(device)).log_softmax(-1)\n    kl_loss = nn.KLDivLoss(reduction=\"batchmean\")\n    loss = kl_loss(logs, y.to(device))\n    model.train()\n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-03-29T21:13:57.721022Z","iopub.execute_input":"2024-03-29T21:13:57.721424Z","iopub.status.idle":"2024-03-29T21:13:57.729747Z","shell.execute_reply.started":"2024-03-29T21:13:57.721391Z","shell.execute_reply":"2024-03-29T21:13:57.728733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@t.no_grad()\ndef augment_data(data, alpha=0.3):\n    # was alpha=0.01 try 0.3\n    # data → ('batch', 'seq', 'channel')\n    data = data.to(device)\n    std = data.std(dim=1, keepdim=True)\n    noise = t.randn_like(data, device=device) * std * alpha\n    return data + noise\n\ndef augment_if(data, iter):\n    if iter == 0: return data\n    return augment_data(data)","metadata":{"execution":{"iopub.status.busy":"2024-03-29T21:14:00.779611Z","iopub.execute_input":"2024-03-29T21:14:00.780275Z","iopub.status.idle":"2024-03-29T21:14:00.787552Z","shell.execute_reply.started":"2024-03-29T21:14:00.780241Z","shell.execute_reply":"2024-03-29T21:14:00.786385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# config","metadata":{}},{"cell_type":"code","source":"batch_size = 55\nprefetch_factor = 10\nnum_workers = 3","metadata":{"execution":{"iopub.status.busy":"2024-03-29T20:45:07.284215Z","iopub.execute_input":"2024-03-29T20:45:07.285132Z","iopub.status.idle":"2024-03-29T20:45:07.289183Z","shell.execute_reply.started":"2024-03-29T20:45:07.285098Z","shell.execute_reply":"2024-03-29T20:45:07.288167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# data","metadata":{}},{"cell_type":"code","source":"test_path = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/'\ntrain_path = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/'\ntrain_spec_path = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/'\nBASE_PATH = '/kaggle/input/hms-harmful-brain-activity-classification/'\n# /kaggle/input/hms-harmful-brain-activity-classification/train.csv\n# PRE_PROCESSED_PATH = './preprocessed/'\n# PRE_PROCESSED_PATH = './eeg-filtered/'\n# PRE_PROCESSED_PATH = './eeg-logged/'\n# PRE_PROCESSED_PATH = './eeg-robust-filter/'\n# PRE_PROCESSED_PATH = './eeg-band-1-70/'\nPRE_PROCESSED_PATH = './eeg-band-1-70-notch-60/'\n\nFEATS_FOR_REAL = ['Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1', 'Fz', 'Cz', 'Pz', 'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG']\n#                   0      1     2     3     4     5     6     7     8     9    10     11    12    13    14    15    16    17    18    19\n# group by semantic groups LP, LL, RP, RR https://raw.githubusercontent.com/cdeotte/Kaggle_Images/main/Jan-2024/montage.png\n# GROUPS = [\n#     ['Fp1', 'F3', 'C3', 'P3', 'O1'],\n#     ['Fp1', 'F7', 'T3', 'T5', 'O1'],\n#     ['Fp2', 'F4', 'C4', 'P4', 'O2'],\n#     ['Fp2', 'F8', 'T4', 'T6', 'O2'],\n# ]\nGROUPS_IDS = [\n    [0, 1, 2, 3, 7],\n    [0, 4, 5, 6, 7],\n    [11, 12, 13, 14, 18],\n    [11, 15, 16, 17, 18],\n    # [8, 9, 10, 19] # TODO: try with leftovers?\n]\n\nGROUPS_IDS2 = [\n    [0, 1, 2, 3, 7],\n    [0, 4, 5, 6, 7],\n    [11, 12, 13, 14, 18],\n    [11, 15, 16, 17, 18],\n    [8, 9, 10],\n]\n\n\n# LEFTOVERS = [8, 9, 10]\n# EKG = [19]\n# TODO: add frequency domain with fourier's transform\n# TODO: add spectrogram to process with conv2d\n# TODO: merge several models together\n# TODO: when submitting round values close to 0 to exactly 0 and rebalance the rest to sum() == 1 for free boost\n\nTARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote','other_vote']","metadata":{"execution":{"iopub.status.busy":"2024-03-29T19:06:20.062259Z","iopub.execute_input":"2024-03-29T19:06:20.062887Z","iopub.status.idle":"2024-03-29T19:06:20.071931Z","shell.execute_reply.started":"2024-03-29T19:06:20.062855Z","shell.execute_reply":"2024-03-29T19:06:20.070945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(f'{BASE_PATH}/train.csv')\ntest_df = pd.read_csv(f'{BASE_PATH}/test.csv')","metadata":{"execution":{"iopub.status.busy":"2024-03-29T20:03:47.809848Z","iopub.execute_input":"2024-03-29T20:03:47.810752Z","iopub.status.idle":"2024-03-29T20:03:47.979640Z","shell.execute_reply.started":"2024-03-29T20:03:47.810718Z","shell.execute_reply":"2024-03-29T20:03:47.978834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"columns = []","metadata":{"execution":{"iopub.status.busy":"2024-03-29T19:15:18.859118Z","iopub.execute_input":"2024-03-29T19:15:18.859764Z","iopub.status.idle":"2024-03-29T19:15:18.863858Z","shell.execute_reply.started":"2024-03-29T19:15:18.859730Z","shell.execute_reply":"2024-03-29T19:15:18.862925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset(Dataset):\n    def __init__(self):\n        super().__init__()\n        self.dataframe = train_df\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    # @lru_cache(maxsize=None)\n    def __getitem__(self, idx): # not preprocessed RIP\n        row = self.dataframe.iloc[idx]\n        spec_id = row['spectrogram_id']\n        parq_path = f'{train_path}{spec_id}.parquet'\n        spec = pd.read_parquet(parq_path)\n        \n#         spectrogram = pd.read_parquet(f'{SPEC_PATH}{row.spectrogram_id}.parquet')\n        spec_offset = int( row.spectrogram_label_offset_seconds )\n        \n#         spec_offset = row['spectrogram_label_offset_seconds']\n#         offset_dp = int(start_time_second * 200)\n#         duration = 10 * 60 * 200\n        columns = spec.columns\n    \n        spec = spec.loc[(spec.time>=spec_offset)\n                             &(spec.time<spec_offset+600)]\n#         spec = spec.iloc[offset_dp:offset_dp+duration]\n        spec = spec.ffill(axis=0)\n        spec = spec.fillna(0)\n        prefixes = [\"LL\",\"RL\",\"RP\",\"LP\"]\n        things = []\n        \n        for i, prefix in enumerate(prefixes):\n            prefix_df = spec.filter(regex=f'^{prefix}', axis=1)\n            things.append(t.tensor(prefix_df.values))\n        samples = torch.stack(things)\n            \n        \n        labels = row[TARGETS].values.astype(np.float64)\n        \n        labels = labels/np.sum(labels)\n#         samples = t.tensor(spec[spec.columns].values)\n        labels_out = t.tensor(labels,dtype=t.float64)\n        \n        assert not samples.isnan().any()\n        assert not labels_out.isnan().any()\n        \n        return samples, labels_out","metadata":{"execution":{"iopub.status.busy":"2024-03-29T20:36:23.157986Z","iopub.execute_input":"2024-03-29T20:36:23.158813Z","iopub.status.idle":"2024-03-29T20:36:23.169504Z","shell.execute_reply.started":"2024-03-29T20:36:23.158779Z","shell.execute_reply":"2024-03-29T20:36:23.168569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def keep_concensus(df):\n    idx = df[TARGETS].max(axis=1) == df[TARGETS].sum(axis=1)\n    return df[idx]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"blah = train_df.groupby(\"spectrogram_id\")\nspec = pd.read_parquet(f'{test_path}853520.parquet')\n\ni = 0\nfor a, b in blah:\n    if i > 5:\n        break\n    print(b[TARGETS])\n    i += 1","metadata":{"execution":{"iopub.status.busy":"2024-03-29T20:11:52.878005Z","iopub.execute_input":"2024-03-29T20:11:52.878337Z","iopub.status.idle":"2024-03-29T20:11:52.937528Z","shell.execute_reply.started":"2024-03-29T20:11:52.878311Z","shell.execute_reply":"2024-03-29T20:11:52.936579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_spectogram(spec_df, prefixes, title = \"Spectogram\"):\n    fig = sp.make_subplots(rows=len(prefixes), cols=1, subplot_titles=prefixes)\n    for i, prefix in enumerate(prefixes):\n        prefix_df = spec_df.filter(regex=f'^{prefix}', axis=1)\n        epsilon = 1e-10\n        fig.add_trace(go.Heatmap(z=np.log(prefix_df + epsilon).T,\n                                 y=pd.to_numeric(prefix_df.columns.str.replace(f\"{prefix}_\", '')),\n                                 coloraxis=\"coloraxis\"),\n                      row=i+1, col=1)\n         # Update x-axis and y-axis labels\n        fig.update_xaxes(title_text=\"Time(Seconds)\", row=i+1, col=1)\n        fig.update_yaxes(title_text=\"Frequency(Hz)\", row=i+1, col=1)\n        # update coloraxis\n        fig.update_layout(coloraxis = {'colorscale':'Jet'}, height=1500,title_text=title)\n    fig.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-29T19:32:52.946745Z","iopub.execute_input":"2024-03-29T19:32:52.947658Z","iopub.status.idle":"2024-03-29T19:32:52.955390Z","shell.execute_reply.started":"2024-03-29T19:32:52.947619Z","shell.execute_reply":"2024-03-29T19:32:52.954371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = Dataset()\n\nsample, labels = test_data[600]\n\nlogified = np.log(sample.permute(0, 2, 1) + 1e-10)\nfor channel in range(4):\n    plt.imshow(logified[channel])\n    print(f\"{sample.shape=}\")\n    plt.show()\n\n\n\n# Size is 4, 300, 100\n\n","metadata":{"execution":{"iopub.status.busy":"2024-03-29T20:37:55.411063Z","iopub.execute_input":"2024-03-29T20:37:55.411434Z","iopub.status.idle":"2024-03-29T20:37:56.476419Z","shell.execute_reply.started":"2024-03-29T20:37:55.411401Z","shell.execute_reply":"2024-03-29T20:37:56.475559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = Dataset()\n# TODO\n# train_df = keep_concensus(train_df)\nids = train_df['spectrogram_id'].unique()\nnp.random.shuffle(ids)\nsplit = int(len(ids) * 0.95)\n\ntrain_ids = ids[:split]\ntest_ids = ids[split:]\n\nnow = datetime.now().strftime(\"%Y-%m-%d_%Hh%M\")\n# t.save(t.tensor(train_ids), f'./splits/{now}_train_ids.pt')\n# t.save(t.tensor(train_ids), f'./splits/{now}_test_ids.pt')\n\ntrain_indices = train_df[train_df['spectrogram_id'].isin(train_ids)].index.tolist()\ntest_indices = train_df[train_df['spectrogram_id'].isin(test_ids)].index.tolist()\n\ntrain_dataset = Subset(dataset, train_indices)\ntest_dataset = Subset(dataset, test_indices)\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=batch_size, num_workers=num_workers, prefetch_factor=prefetch_factor, shuffle=True)\ntest_dataloader = DataLoader(test_dataset, batch_size=batch_size, num_workers=num_workers, prefetch_factor=prefetch_factor, shuffle=True)\n\nlen(train_dataset), len(test_dataset)","metadata":{"execution":{"iopub.status.busy":"2024-03-29T20:45:11.989709Z","iopub.execute_input":"2024-03-29T20:45:11.990547Z","iopub.status.idle":"2024-03-29T20:45:12.016792Z","shell.execute_reply.started":"2024-03-29T20:45:11.990515Z","shell.execute_reply":"2024-03-29T20:45:12.015864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# model 👯‍♀️","metadata":{}},{"cell_type":"markdown","source":"## conv1d + GRU","metadata":{}},{"cell_type":"code","source":"class CNNBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, stride):\n        super(CNNBlock, self).__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(\n                in_channels, out_channels, 4, stride, 1, bias=False, padding_mode=\"reflect\"\n            ),\n            nn.BatchNorm2d(out_channels),\n            nn.LeakyReLU(0.2),\n            nn.Dropout(0.2),\n        )\n\n    def forward(self, x):\n        return self.conv(x)\n        \nclass Model(nn.Module):\n    def __init__(self, in_channels=4, features=[i**2 for i in range(4, 16) if i % 2 == 0], dropout=0.4):\n        super().__init__()\n        layers = [\n            nn.Conv2d(\n                in_channels,\n                features[0],\n                kernel_size=4,\n                stride=2,\n                padding=1,\n                padding_mode=\"reflect\",\n            ),\n            nn.LeakyReLU(0.2),\n        ]\n        in_channels = features[0]\n        \n        for feature in features[1:]:\n            layers.append(\n                CNNBlock(in_channels, feature,\n                         stride=1 if feature == features[-1] else 2),\n            )\n            in_channels = feature\n\n        layers.extend(\n            [\n#                 nn.Conv2d(\n#                     in_channels, feature, kernel_size=4, stride=1, padding=1, padding_mode=\"reflect\"\n#                 ),\n#                 nn.ReLU(),\n                nn.Flatten(start_dim=1, end_dim=-1),\n                nn.Linear(3136, 256),\n                nn.ReLU(),\n                nn.Dropout(dropout),\n                nn.Linear(256, 6),\n            ]\n        )\n\n        self.model = nn.Sequential(*layers)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x\n\ndef test():\n    m = Model().to(device)\n    x, y = next(train_dataloader.__iter__())\n    print(f'{x.shape=}')\n    r = m(x.to(device))\n    print(f'{r.shape=}')\n    \ntest()","metadata":{"execution":{"iopub.status.busy":"2024-03-29T21:02:35.320217Z","iopub.execute_input":"2024-03-29T21:02:35.320633Z","iopub.status.idle":"2024-03-29T21:02:40.982361Z","shell.execute_reply.started":"2024-03-29T21:02:35.320600Z","shell.execute_reply":"2024-03-29T21:02:40.981275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## transformer","metadata":{}},{"cell_type":"code","source":"class Transformer(nn.Module):\n    def __init__(self, d_chan=20, d_model=256, d_clump=4):\n        super().__init__()\n        self.d_clump = d_clump\n\n        self.start = nn.Parameter(t.randn(1, 1, d_model))\n        self.bn = nn.BatchNorm1d(d_chan)\n        self.emb = nn.Linear(d_chan * d_clump, d_model)\n        self.llm = nn.Transformer(d_model=d_model, nhead=8, num_encoder_layers=3, num_decoder_layers=0, dim_feedforward=d_model * 2, batch_first=True)\n        self.head = nn.Sequential(\n            nn.Linear(d_model, d_model // 2),\n            nn.ReLU(),\n            nn.Linear(d_model // 2, 6)\n        )\n\n    def forward(self, x):\n        x = self.bn(x.permute((0, 2, 1))).permute((0, 2, 1))\n        x = einops.rearrange(x, 'batch (seq clump) channels -> batch seq (clump channels)', clump=self.d_clump)\n        x = self.emb(x)\n        # add a fake start token\n        x = t.cat([self.start.repeat(x.shape[0], 1, 1), x], dim=1)\n        x = self.llm.encoder(x)[:, 0]\n        return self.head(x)\n\ndef scope():\n    val, label = next(train_dataloader.__iter__())\n    model = Transformer().to(device)\n    output = model(val.to(device))\n\n# scope()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## separated GRU","metadata":{}},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self, d_in, d_out, kernel_size, drop):\n        super().__init__()\n        self.model = nn.Sequential(\n            nn.Conv1d(d_in, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Conv1d(d_out, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Conv1d(d_out, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.MaxPool1d(kernel_size=2, stride=2, padding=0), # reduce sequence size by 2\n        )\n    def forward(self, x):\n        # TODO: add skip for training speed\n        return self.model(x)\n        \nclass SeparatedGRU_old(nn.Module):\n    def __init__(self, in_channels=5, gru_hidden_size=128, drop=0.2):\n        super().__init__()\n        self.d_split = len(GROUPS_IDS)\n        self.pre_out = in_channels * 4\n        self.gru_hidden_size = gru_hidden_size\n        \n        self.pre_process = nn.Sequential(\n            # nn.LayerNorm(normalized_shape=[in_channels, 10000]),\n            nn.BatchNorm1d(in_channels, momentum=None),\n            # use conv1d as a denoiser\n            # block 1\n            ConvBlock(in_channels, in_channels * 2, kernel_size=3, drop=drop),\n            # nn.BatchNorm1d(in_channels * 2, momentum=None),\n            \n            # block 2\n            ConvBlock(in_channels * 2, in_channels * 4, kernel_size=5, drop=drop),\n            # nn.BatchNorm1d(in_channels * 4, momentum=None),\n\n            # block 3\n            ConvBlock(in_channels * 4, in_channels * 4, kernel_size=7, drop=drop),\n            # nn.BatchNorm1d(in_channels * 4, momentum=None), # re-enable one more to force training ?\n        )\n        \n        # TODO: add a learnable first state for GRU or check what is the default\n        self.gru = nn.GRU(self.pre_out, self.gru_hidden_size, num_layers=1, batch_first=True, bidirectional=True)\n\n        self.post_gru = nn.Sequential(\n            # nn.BatchNorm1d(self.gru_hidden_size * 2, momentum=None),\n            nn.Dropout(drop),\n            nn.Linear(self.gru_hidden_size * 2, self.gru_hidden_size),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Linear(self.gru_hidden_size, self.gru_hidden_size),\n        )\n\n        self.head = nn.Sequential(\n            # nn.BatchNorm1d(self.gru_hidden_size * self.d_split, momentum=None),\n            nn.Linear(self.gru_hidden_size * self.d_split, self.gru_hidden_size * 2),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Linear(self.gru_hidden_size * 2, 6)\n        )\n\n    def forward(self, x: ('batch', 'seq', 'channel')):\n        # separate the input into 4 splits (LP, LL, RP, RR)\n        splits = [x[:, :, group] for group in GROUPS_IDS]\n        # fold it into batch so we can run in parallel\n        x = einops.rearrange(t.stack(splits, dim=0), 'group batch seq channel -> (group batch) seq channel')\n\n        # pre_process: (batch, channel, seq) → (batch / 4, channel * 4, seq)\n        x = x.permute((0, 2, 1))\n        x = self.pre_process(x)\n        x = x.permute((0, 2, 1))\n\n        # GRU: (batch, seq, input_size), [(2 * num_layers, batch, hidden_size)] → (batch, seq, 2 * hidden_size)\n        x, _ = self.gru(x)\n        x = x[:, -1, :]\n\n        # MLP post GRU\n        x = self.post_gru(x)\n\n        # unfold the splits\n        x = einops.rearrange(x, '(group batch) hidden -> batch (hidden group)', group=self.d_split)\n\n        # head: (batch, 2 * hidden_size) → (batch, 6)\n        x = self.head(x)\n\n        # out: → (batch, 6)\n        return x\n\ndef scope():\n    m = SeparatedGRU().to(device)\n    x, y = next(train_dataloader.__iter__())\n    r = m(x.to(device))\n    print(f'{r.shape=}')\n    \n# scope()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## separated GRU w/ montage","metadata":{}},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self, d_in, d_out, kernel_size, drop):\n        super().__init__()\n        self.model = nn.Sequential(\n            nn.Conv1d(d_in, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.InstanceNorm1d(d_out),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Conv1d(d_out, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            # nn.InstanceNorm1d(d_out),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Conv1d(d_out, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            # nn.InstanceNorm1d(d_out),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.MaxPool1d(kernel_size=2, stride=2, padding=0), # reduce sequence size by 2\n        )\n    def forward(self, x):\n        # TODO: add skip for training speed\n        return self.model(x)\n        \nclass SeparatedGRU(nn.Module):\n    def __init__(self, in_channels=19, gru_hidden_size=128, drop=0.3):\n    # def __init__(self, in_channels=4, gru_hidden_size=256, drop=0.5):\n        super().__init__()\n        # TODO double the size of channels, see if it helps?\n        in_channels *= 2\n        self.d_split = len(GROUPS_IDS)\n        self.pre_out = in_channels * 4\n        self.gru_hidden_size = gru_hidden_size\n        \n        self.pre_process = nn.Sequential(\n            # nn.InstanceNorm1d(in_channels),\n            # nn.LayerNorm(normalized_shape=[in_channels, 10000]),\n            # nn.BatchNorm1d(in_channels, momentum=None),\n            # use conv1d as a denoiser\n            # -- block 1 --\n            ConvBlock(in_channels // 2, in_channels * 2, kernel_size=3, drop=drop),\n            # nn.BatchNorm1d(in_channels * 2, momentum=None),\n            # nn.LayerNorm(normalized_shape=[in_channels * 2, 5000]),\n            # -- block 2 --\n            ConvBlock(in_channels * 2, in_channels * 4, kernel_size=5, drop=drop),\n            # nn.BatchNorm1d(in_channels * 4, momentum=None),\n            # nn.LayerNorm(normalized_shape=[in_channels * 4, 2500]),\n            # -- block 3 --\n            ConvBlock(in_channels * 4, in_channels * 4, kernel_size=7, drop=drop),\n            # nn.BatchNorm1d(in_channels * 4, momentum=None), # re-enable one more to force training ?\n            # nn.LayerNorm(normalized_shape=[in_channels * 4, 1250]),\n        )\n        \n        # TODO: add a learnable first state for GRU or check what is the default\n        self.gru = nn.GRU(self.pre_out, self.gru_hidden_size, num_layers=1, batch_first=True, bidirectional=True)\n\n        self.post_gru = nn.Sequential(\n            # nn.BatchNorm1d(self.gru_hidden_size * 2, momentum=None),\n            nn.LayerNorm(normalized_shape=[self.gru_hidden_size * 2]),\n            nn.Dropout(drop),\n            # nn.LayerNorm(normalized_shape=self.gru_hidden_size * 2),\n            nn.Linear(self.gru_hidden_size * 2, self.gru_hidden_size),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            # nn.LayerNorm(normalized_shape=self.gru_hidden_size),\n            nn.Linear(self.gru_hidden_size, self.gru_hidden_size),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            # nn.LayerNorm(normalized_shape=self.gru_hidden_size),\n            nn.Linear(self.gru_hidden_size, 6),\n        )\n\n        # self.head = nn.Sequential(\n        #     # nn.BatchNorm1d(self.gru_hidden_size * self.d_split, momentum=None),\n        #     # nn.LayerNorm(normalized_shape=[self.gru_hidden_size * self.d_split]),\n        #     nn.Linear(self.gru_hidden_size * self.d_split, self.gru_hidden_size * 2),\n        #     nn.ReLU(),\n        #     nn.Dropout(drop),\n        #     nn.Linear(self.gru_hidden_size * 2, 6)\n        # )\n\n    def montage(self, x):\n        splits = [x[:, :, group] for group in GROUPS_IDS]\n        splits = [s[:, :, :-1] - s[:, :, 1:] for s in splits]\n        return einops.rearrange(t.stack(splits, dim=0), 'group batch seq channel -> (group batch) seq channel')\n    \n    def montage2(self, x):\n        splits = [x[:, :, group] for group in GROUPS_IDS2]\n        splits = [s[:, :, :-1] - s[:, :, 1:] for s in splits]\n        splits.append(x[:, :, 19:20]) # add EKG\n        return t.cat(splits, dim=2)\n        # return einops.rearrange(t.cat(splits, dim=2), 'group batch seq channel -> group seq (batch channel)')\n\n    def forward(self, x: ('batch', 'seq', 'channel')):\n        # separate the input into 4 montages (LP, LL, RP, RR)\n        x = self.montage2(x)\n        # pre_process: (batch, channel, seq) → (batch / 4, channel * 4, seq)\n        x = x.permute((0, 2, 1))\n        x = self.pre_process(x)\n        x = x.permute((0, 2, 1))\n        # GRU: (batch, seq, input_size), [(2 * num_layers, batch, hidden_size)] → (batch, seq, 2 * hidden_size)\n        x, _ = self.gru(x)\n        x = x[:, -1, :]\n        # MLP post GRU\n        x = self.post_gru(x)\n        # unfold the splits\n        # x = einops.rearrange(x, '(group batch) hidden -> batch (hidden group)', group=self.d_split)\n        # head: (batch, 2 * hidden_size) → (batch, 6)\n        # x = self.head(x)\n        # out: → (batch, 6)\n        return x\n\ndef scope():\n    m = SeparatedGRU().to(device)\n    x, y = next(train_dataloader.__iter__())\n    r = m(x.to(device))\n    print(f'{r.shape=}')\n    \n# scope()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train","metadata":{}},{"cell_type":"code","source":"GC()\n# model = Model().to(device)\n# model = Transformer().to(device)\nmodel = Model().to(device)\n# TODO: try cranking the weight decay\n# TODO: try using a scheduler\nopt = t.optim.Adam(model.parameters(), lr=3e-4, weight_decay=1e-5)\nprint(f'model has {sum(p.numel() for p in model.parameters())} params')","metadata":{"execution":{"iopub.status.busy":"2024-03-29T21:04:08.984919Z","iopub.execute_input":"2024-03-29T21:04:08.985271Z","iopub.status.idle":"2024-03-29T21:04:11.427419Z","shell.execute_reply.started":"2024-03-29T21:04:08.985245Z","shell.execute_reply":"2024-03-29T21:04:11.426486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## base","metadata":{}},{"cell_type":"code","source":"def train(model, opt, wnb=True, data_augmentation=False):\n    model.train()\n    validation_test, validation_test_label = next(test_dataloader.__iter__())\n    validation_train, validation_train_label = next(train_dataloader.__iter__())\n\n    limited_replay_buffer = []\n    tq = tqdm(train_dataloader)\n    for i, (x_train, y_train) in enumerate(tq):        \n        limited_replay_buffer.append((x_train, y_train))\n        if i > 1: break\n\n    if wnb:\n        wandb.init(project='kaggle-eeg-rc')\n        now = datetime.now().strftime(\"%Y-%m-%d_%Hh%M\")\n        wandb.log({'val_test':   eval(model, validation_test, validation_test_label, do_eval=True), 'now': f'{now}'})\n        wandb.log({'val_train':  eval(model, validation_train, validation_train_label, do_eval=True), 'now': f'{now}'})\n    for epoch in range(1000):\n        replay_buffer, maxi = [], 3\n        # tq = tqdm(train_dataloader)\n        # for x_train, y_train in tq:\n        for _ in range(1):\n            # replay_buffer.append((x_train, y_train))\n            # replay_buffer = replay_buffer[-maxi:]\n            # XXX\n            replay_buffer = limited_replay_buffer\n            for x_train, y_train in replay_buffer: # burn more GPU otherwise we bottleneck on disk IO\n            # for k in range(3): # the data reading is too slow, so force the GPU to spin\n                if data_augmentation: x_train = augment_data(x_train)\n                logs = model(x_train.to(device)).log_softmax(-1)\n                kl_loss = nn.KLDivLoss(reduction=\"batchmean\")\n                loss = kl_loss(logs, y_train.to(device))\n                opt.zero_grad()\n                loss.backward()\n                opt.step()\n                # tq.set_description(f'loss = {loss:.4f}')\n                if wnb: wandb.log({'loss': loss.item()})\n        \n        now = datetime.now().strftime(\"%Y-%m-%d_%Hh%M\")\n        if wnb:\n            wandb.log({'val_test':   eval(model, validation_test, validation_test_label, do_eval=True), 'now': f'{now}'})\n            wandb.log({'val_train':  eval(model, validation_train, validation_train_label, do_eval=True), 'now': f'{now}'})\n        t.save(model.state_dict(), f'weights/gru-4-splits_{now}.pt')\n    if wnb: wandb.finish()\n\n# train(model, opt, wnb=True, data_augmentation=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-29T21:04:14.410303Z","iopub.execute_input":"2024-03-29T21:04:14.411770Z","iopub.status.idle":"2024-03-29T21:04:14.427239Z","shell.execute_reply.started":"2024-03-29T21:04:14.411720Z","shell.execute_reply":"2024-03-29T21:04:14.426053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## staged","metadata":{}},{"cell_type":"code","source":"def populate_buffer(buffer_size=300):\n    while True:\n        limited_replay_buffer = []\n        tq = tqdm(train_dataloader)\n        for x_train, y_train in tq:\n            limited_replay_buffer.append((x_train, y_train))\n            if len(limited_replay_buffer) >= buffer_size:\n                yield limited_replay_buffer\n                del limited_replay_buffer\n                limited_replay_buffer = []\n\ndef staged_train(model, opt, mini_epochs, wnb=True, data_augmentation=False):\n    model.train()\n    validation_test, validation_test_label = next(test_dataloader.__iter__())\n    validation_train, validation_train_label = next(train_dataloader.__iter__())\n\n    if wnb:\n        wandb.init(project='kaggle-eeg-rc')\n        now = datetime.now().strftime(\"%Y-%m-%d_%Hh%M\")\n        wandb.log({'val_test':   eval(model, validation_test, validation_test_label, do_eval=True), 'now': f'{now}'})\n        wandb.log({'val_train':  eval(model, validation_train, validation_train_label, do_eval=True), 'now': f'{now}'})\n\n    for replay_buffer in populate_buffer():\n        for epoch in tqdm(range(mini_epochs)):\n            for x_train, y_train in replay_buffer:\n                # TODO: use variable alpha based on mini_epoch\n                # if data_augmentation: x_train = augment_data(x_train, alpha=epoch / 10)\n                if data_augmentation: x_train = augment_data(x_train, alpha=random.choice([0.01, 0.05, 0.1])) #, 0.15, 0.2, 0.3]))\n                # if data_augmentation: x_train = augment_data(x_train, alpha=0.2)\n                logs = model(x_train.to(device)).log_softmax(-1)\n                kl_loss = nn.KLDivLoss(reduction=\"batchmean\")\n                loss = kl_loss(logs, y_train.to(device))\n                opt.zero_grad()\n                loss.backward()\n                opt.step()\n                if wnb: wandb.log({'loss': loss.item()})\n        \n            now = datetime.now().strftime(\"%Y-%m-%d_%Hh%M\")\n            if wnb:\n                wandb.log({'val_test':   eval(model, validation_test, validation_test_label, do_eval=True), 'now': f'{now}'})\n                wandb.log({'val_train':  eval(model, validation_train, validation_train_label, do_eval=True), 'now': f'{now}'})\n            t.save(model.state_dict(), f'weights/gru-4-splits_{now}.pt')\n        del replay_buffer\n    if wnb: wandb.finish()\n\nstaged_train(model, opt, mini_epochs=20, wnb=True, data_augmentation=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-29T21:14:10.390404Z","iopub.execute_input":"2024-03-29T21:14:10.391136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# save / load","metadata":{}},{"cell_type":"code","source":"# t.save(model.state_dict(),'model-weights4.pt')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = SeparatedGRU().to(device)\n# model.load_state_dict(t.load('weights/gru-4-splits_2024-03-24_15h03.pt', map_location=device))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_train, y_train = next(train_dataloader.__iter__())\nx_val, y_val = next(test_dataloader.__iter__())\n\nprint(f'train(): train: {eval(model, x_train, y_train, do_eval=False)}')\nprint(f'train(): test:  {eval(model, x_val, y_val, do_eval=False)}')\nprint('--')\nprint(f'eval(): train   {eval(model, x_train, y_train, do_eval=True)}')\nprint(f'eval(): test    {eval(model, x_val, y_val, do_eval=True)}')\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# submit","metadata":{}},{"cell_type":"code","source":"@t.no_grad()\ndef submit(model, test_dataloader, test_df):\n    model.eval()\n    res = []\n    for batch in test_dataloader:\n        prob = model(batch.to(device)).softmax(-1)\n        res.append(prob.detach().cpu())\n\n    res = t.cat(res, dim=0)\n    sub = test_df[[\"eeg_id\"]].copy()\n    sub[TARGETS] = res\n    sub.to_csv('submission.csv',index=False)\n    print('Submission shape',sub.shape)\n    display(sub.head())","metadata":{},"execution_count":null,"outputs":[]}]}