{"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":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7615793,"sourceType":"datasetVersion","datasetId":4433330},{"sourceId":7884584,"sourceType":"datasetVersion","datasetId":4628215}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## **NNSPT 1D CNN Starter using RAW EEG Features!**","metadata":{}},{"cell_type":"markdown","source":"## **Install deps**","metadata":{}},{"cell_type":"code","source":"!pip install ../input/nnsptwheel/nnspt-0.0.2-py2.py3-none-any.whl --no-deps &> null\n!pip install ../input/ecgmentationswheel/ecgmentations-0.0.7-py2.py3-none-any.whl --no-deps &> null","metadata":{"execution":{"iopub.status.busy":"2024-03-18T08:00:30.375681Z","iopub.execute_input":"2024-03-18T08:00:30.376054Z","iopub.status.idle":"2024-03-18T08:00:32.595243Z","shell.execute_reply.started":"2024-03-18T08:00:30.376021Z","shell.execute_reply":"2024-03-18T08:00:32.593637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Load data**","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport polars as pd\n\nfrom tqdm import tqdm\nfrom pathlib import Path\nfrom collections import Counter\nfrom scipy.signal import butter, iirnotch, lfilter\n\nseed = 1996\nnfolds = 5\n\nhz = 200\nduration = 50\n\npath = Path('../input/hms-harmful-brain-activity-classification')\n\nclasses = [ 'seizure_vote',\n            'lpd_vote',\n            'gpd_vote',\n            'lrda_vote', \n            'grda_vote',\n            'other_vote' ]\n\nleads = [ 'Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1', 'Fz', 'Cz', 'Pz', 'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG' ]\nleads2idx = { lead: idx for idx, lead in enumerate(leads) }\n\nfleads = [ 'Fp1', 'F3', 'C3', 'P3',\n           'Fp2', 'F4', 'C4', 'P4',\n           'Fp1', 'F7', 'T3', 'T5',\n           'Fp2', 'F8', 'T4', 'T6',\n            'Fz', 'Cz',\n         ]\n\nsleads = [ 'F3', 'C3', 'P3', 'O1',\n           'F4', 'C4', 'P4', 'O2',\n           'F7', 'T3', 'T5', 'O1',\n           'F8', 'T4', 'T6', 'O2',\n           'Cz', 'Pz',\n         ]\n\neidx = leads2idx['EKG']\n\nfids = [ leads2idx[lead] for lead in fleads]\nsids = [ leads2idx[lead] for lead in sleads]\n\ntrain_eeg_path = path / 'train_eegs'\ntest_eeg_path = path / 'test_eegs'\n\ntrain = pd.read_csv(path / 'train.csv')\n\ntrain_data = {}\n\neegvalue = 200\necgvalue = 3\n\nfor idx in tqdm([*train['eeg_id'].unique()]):\n    data = pd.read_parquet(train_eeg_path / f'{idx}.parquet').to_numpy()\n\n    # remove powerline noise 60 Hz\n    b, a = butter(4, (59, 61), 'bandstop', analog=False, fs=200)\n    data = lfilter(b, a, data, axis=0)\n\n    eeg = data[:, fids] - data[:, sids]\n    \n    # remove extra freq eeg\n\n    b, a = butter(4, (1, 70), 'bandpass', analog=False, fs=200)\n    eeg = lfilter(b, a, eeg, axis=0)\n\n    # make normalization eeg\n\n    eeg = np.clip(eeg, -eegvalue, eegvalue)\n    eeg = np.nan_to_num(eeg, nan=0) / eegvalue\n\n    # mu coding eeg\n\n    mu = 1.\n    eeg = np.sign(eeg) * np.log(1 + mu * np.abs(eeg)) / np.log(mu + 1)\n\n    # make normalization ecg\n\n    ecg = data[:, eidx:eidx+1]\n    ecg = np.nan_to_num(ecg, nan=0) / 1000\n    ecg = np.clip(ecg, -ecgvalue, ecgvalue) / ecgvalue\n    \n    data = np.concatenate([eeg, ecg], axis=-1)\n\n    train_data[idx] = data.astype(np.float16)\n\nhigh_quality_data_info = {}\n\nfor (idx, sidx), group in tqdm(train.group_by(['eeg_id', 'eeg_sub_id'])):\n    key = f'{idx}'\n\n    patient = group['patient_id'][0]\n    consensus = group['expert_consensus'][0]\n\n    offset = int(group['eeg_label_offset_seconds'][0] * hz)\n    label = group[classes].to_numpy()[0]\n\n    nvoters = label.sum()\n\n    if nvoters > 5:\n        if key not in high_quality_data_info:\n            high_quality_data_info[key] = {\n                'eidx': idx,\n                'offsets': [ offset ],\n                'labels': [ label],\n                'consensus': [ consensus ],\n                'nvoters': [ nvoters ],\n                'patient': patient,\n                'nsubrecords': 1\n            }\n        else:\n            high_quality_data_info[key]['offsets'].append(offset)\n            high_quality_data_info[key]['labels'].append(label)\n            high_quality_data_info[key]['consensus'].append(consensus)\n            high_quality_data_info[key]['nvoters'].append(nvoters)\n            high_quality_data_info[key]['nsubrecords'] += 1\n\nkeys = list()\n\nfor key in high_quality_data_info:\n    counts =  Counter(high_quality_data_info[key]['consensus']).most_common()\n\n    if len(counts) > 3:\n        if counts[0][1] == counts[1][1] == counts[2][1]:\n            cons = [counts[0][0], counts[1][0], counts[2][0]]\n            cons = sorted(cons)\n\n            consensus = cons[0]\n            continue\n        elif counts[0][1] == counts[1][1]:\n            cons = [counts[0][0], counts[1][0]]\n            cons = sorted(cons)\n\n            consensus = cons[0]\n            continue\n        else:\n            consensus = counts[0][0]\n    elif len(counts) == 2:\n        if counts[0][1] == counts[1][1]:\n            cons = [counts[0][0], counts[1][0]]\n            cons = sorted(cons)\n\n            consensus = cons[0]\n            continue\n        else:\n            consensus = counts[0][0]\n    else:\n        consensus = counts[0][0]\n\n    high_quality_data_info[key]['consensus'] = consensus\n\n    keys.append(key)\n\nkeys = sorted(keys)\n\nconsensus = [high_quality_data_info[key]['consensus'] for key in keys ]\npatients = [high_quality_data_info[key]['patient'] for key in keys ]\n\ntest_data = {}\n\ntest = pd.read_csv(path / 'test.csv')\n\nfor idx in tqdm([*test['eeg_id'].unique()]):\n    data = pd.read_parquet(test_eeg_path / f'{idx}.parquet').to_numpy()\n\n    # remove powerline noise 60 Hz\n    b, a = butter(4, (59, 61), 'bandstop', analog=False, fs=200)\n    data = lfilter(b, a, data, axis=0)\n\n    eeg = data[:, fids] - data[:, sids]\n    \n    # remove extra freq eeg\n\n    b, a = butter(4, (1, 70), 'bandpass', analog=False, fs=200)\n    eeg = lfilter(b, a, eeg, axis=0)\n\n    # make normalization eeg\n\n    eeg = np.clip(eeg, -eegvalue, eegvalue)\n    eeg = np.nan_to_num(eeg, nan=0) / eegvalue\n\n    # mu coding eeg\n\n    mu = 1.\n    eeg = np.sign(eeg) * np.log(1 + mu * np.abs(eeg)) / np.log(mu + 1)\n\n    # make normalization ecg\n\n    ecg = data[:, eidx:eidx+1]\n    ecg = np.nan_to_num(ecg, nan=0) / 1000\n    ecg = np.clip(ecg, -ecgvalue, ecgvalue) / ecgvalue\n    \n    data = np.concatenate([eeg, ecg], axis=-1)\n\n    test_data[idx] = data.astype(np.float16)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T08:01:05.488029Z","iopub.execute_input":"2024-03-18T08:01:05.488853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Split data**","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedGroupKFold\n\nfolds = StratifiedGroupKFold(n_splits=nfolds, shuffle=True, random_state=seed)\nsplits = folds.split(keys, consensus, patients)\n\nsplits = iter(splits)\nsplit = next(splits)\n\ntrain_ids, valid_ids = split","metadata":{"execution":{"iopub.status.busy":"2024-03-18T08:00:35.449849Z","iopub.status.idle":"2024-03-18T08:00:35.450718Z","shell.execute_reply.started":"2024-03-18T08:00:35.450437Z","shell.execute_reply":"2024-03-18T08:00:35.450466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Create dataloaders**","metadata":{}},{"cell_type":"code","source":"import ecgmentations as E\n\nfrom collections import Counter\nfrom torch.utils.data import Dataset, DataLoader, Sampler\n\nclass EEGDataset(Dataset):\n    def __init__(self, info, data, keys, ids, augs=dict):\n        self.info = info\n        self.data = data\n        self.augs = augs\n\n        self.keys = keys\n        self.ids = ids\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, idx):\n        idx = self.ids[idx]\n\n        key = self.keys[idx]\n        info = self.info[key]\n\n        sidx = np.random.choice(info['nsubrecords'])\n        offset = info['offsets'][sidx]\n\n        eeg = self.data[info['eidx']][offset:offset+duration*hz]\n        label = info['labels'][sidx]\n        nvoters = info['nvoters'][sidx]\n        \n        label = label / nvoters\n\n        return np.ascontiguousarray(eeg.T).astype(np.float32), label\n\nclass Sampler(Sampler):\n    def __init__(self, dataset):\n        self.dataset = dataset\n\n        counts = Counter([ dataset.info[dataset.keys[idx]]['consensus'] for idx in dataset.ids ])\n\n        cprobs = {\n            'Other':   0.500,\n            'LPD':     0.150,\n            'GPD':     0.150,\n            'GRDA':    0.080,\n            'LRDA':    0.080,\n            'Seizure': 0.040\n        }\n\n        self.probs = []\n\n        for idx in dataset.ids:\n            consensus = dataset.info[dataset.keys[idx]]['consensus']\n\n            prob = cprobs[consensus] / counts[consensus]\n            self.probs.append(prob)\n\n        self.size = len(dataset)\n\n    def __len__(self):\n        return self.size\n\n    def __iter__(self):\n        for _ in range(self.size):\n            idx = np.random.choice(self.size, p=self.probs)\n            yield idx\n    \nclass LeftRightSwap(E.EcgOnlyTransform):\n    def apply(self, ecg, **params):\n        ecg[:, [0, 1, 2, 3, 4, 5, 6, 7]] = ecg[:, [4, 5, 6, 7, 0, 1, 2, 3]]        \n        ecg[:, [8, 9, 10, 11, 12, 13, 14, 15]] = ecg[:, [12, 13, 14, 15, 8, 9, 10, 11, ]]        \n\n        return ecg\n\n    def get_transform_init_args_names(self):\n        return tuple()\n\naugs = E.Sequential([\n    LeftRightSwap(p=0.5),\n    E.AmplitudeInvert(p=0.25),\n    E.RandomTimeWrap(p=0.1),\n    E.TimeReverse(p=0.1),\n    E.TimeShift(p=0.75),\n    E.TimeSegmentShuffle(p=0.15),\n])\n\ntrain_dataset = EEGDataset(high_quality_data_info, train_data, keys, train_ids, augs=augs)\nsampler = Sampler(dataset=train_dataset)\n\ntrain_dataloader = DataLoader(\n    train_dataset,\n    32,\n    shuffle=False,\n    num_workers=8,\n    prefetch_factor=4,\n    sampler=sampler,\n)\n\nvalid_dataset = EEGDataset(high_quality_data_info, train_data, keys, valid_ids)\nvalid_dataloader = DataLoader(valid_dataset, 32, shuffle=False, num_workers=4, prefetch_factor=4)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T08:00:35.452004Z","iopub.status.idle":"2024-03-18T08:00:35.452937Z","shell.execute_reply.started":"2024-03-18T08:00:35.452693Z","shell.execute_reply":"2024-03-18T08:00:35.452713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Create model**","metadata":{}},{"cell_type":"code","source":"import torch\n\nfrom nnspt.blocks.encoders import Encoder\n\nclass ParallelConv1d(torch.nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        stride,\n    ):\n        super().__init__()\n\n        self.conv_3 = torch.nn.Conv1d(in_channels, out_channels, kernel_size=(3,), stride=(stride,), padding=(1,), bias=False)\n        self.conv_5 = torch.nn.Conv1d(in_channels, out_channels, kernel_size=(5,), stride=(stride,), padding=(2,), bias=False)\n        self.conv_7 = torch.nn.Conv1d(in_channels, out_channels, kernel_size=(7,), stride=(stride,), padding=(3,), bias=False)\n        self.conv_9 = torch.nn.Conv1d(in_channels, out_channels, kernel_size=(9,), stride=(stride,), padding=(4,), bias=False)\n\n    def forward(self, x):\n        x3 = self.conv_3(x)\n        x5 = self.conv_5(x)\n        x7 = self.conv_7(x)\n        x9 = self.conv_9(x)\n\n        x = x3 + x5 + x7 + x9\n\n        return x\n\nclass SimpleClassifier(torch.nn.Module):\n    def __init__(self, nleads, encoder, p=0.15):\n        super().__init__()\n\n        self.encoder = Encoder(in_channels=nleads, depth=5, name=encoder)\n        self.encoder.conv_stem = ParallelConv1d(nleads, 24, 2)\n        \n        self.other = torch.nn.Sequential(\n            torch.nn.AdaptiveAvgPool1d(1),\n            torch.nn.Flatten(),\n            torch.nn.Dropout(p=p, inplace=True),\n            torch.nn.Linear(self.encoder.out_channels[-4], 2)\n        )\n        \n        self.head = torch.nn.Sequential(\n            torch.nn.AdaptiveAvgPool1d(1),\n            torch.nn.Flatten(),\n            torch.nn.Dropout(p=p, inplace=True),\n            torch.nn.Linear(self.encoder.out_channels[-1], 6)\n        )\n\n    def forward(self, x):\n        f = self.encoder(x)\n\n        z = self.other(f[-4])\n        x = self.head(f[-1])\n\n        return x, z","metadata":{"execution":{"iopub.status.busy":"2024-03-18T08:00:35.454165Z","iopub.status.idle":"2024-03-18T08:00:35.454649Z","shell.execute_reply.started":"2024-03-18T08:00:35.454389Z","shell.execute_reply":"2024-03-18T08:00:35.454407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Train loop**","metadata":{}},{"cell_type":"code","source":"import copy\nimport torch.optim.swa_utils as tsu\n\nfrom tqdm import tqdm\nfrom sklearn import metrics\n\nnepochs = 60\nsepoch = 1\n\neps = 1e-15\n\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n\nmodel = SimpleClassifier(nleads=19, encoder='timm-efficientnetv2-s')\nmodel.to(device)\n\ndef get_ema_avg_fn(decay=0.99):\n    @torch.no_grad()\n    def ema_update(ema_param, current_param, num_averaged):\n        return decay * ema_param + (1 - decay) * current_param\n\n    return ema_update\n\naveraged_model = tsu.AveragedModel(model, avg_fn=get_ema_avg_fn())\naveraged_model.to(device)\n\n\nopt = torch.optim.AdamW(model.parameters(), lr=0.0005)\nshed = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=0.001, total_steps=nepochs*len(train_dataloader))\n\nloss_his, train_loss = [], []\n\nbest_score = np.inf\nbest_state_dict = copy.deepcopy(averaged_model.state_dict())\n\nfor epoch in range(nepochs):\n    model.train()\n\n    for leads_batch, target_batch in tqdm(train_dataloader):\n        leads_batch = leads_batch.to(device).float()\n        target_batch = target_batch.to(device).float()\n        \n        target_aux_batch = torch.stack([\n            torch.sum( target_batch[:, :5], dim=1),\n            target_batch[:, 5]\n        ], dim=1 )\n\n        pred_logit_batch, pred_aux_logit_batch = model(leads_batch)\n\n        pred_log_softmax_batch = pred_logit_batch.log_softmax(dim=1)\n        loss1 = torch.nn.functional.kl_div(pred_log_softmax_batch, target_batch, reduction='batchmean')\n\n        pred_aux_log_softmax_batch = pred_aux_logit_batch.log_softmax(dim=1)\n        loss2 = torch.nn.functional.kl_div(pred_aux_log_softmax_batch, target_aux_batch, reduction='batchmean')\n\n        loss = loss1 + loss2\n        loss.backward()\n        \n        opt.step()\n        shed.step()\n        opt.zero_grad()\n\n        train_loss.append(loss.item())\n\n        assert not np.isnan(train_loss[-1])\n        \n        averaged_model.update_parameters(model)\n        \n    tsu.update_bn(train_dataloader, averaged_model, device=device)\n\n    loss_his.append(np.mean(train_loss))\n    train_loss.clear()\n\n    print('[Epoch {}/{}] [Loss: {}]'.format(epoch+1, nepochs, loss_his[-1]))\n\n    if (epoch + 1) % sepoch == 0:\n        with torch.no_grad():\n            model.eval()\n\n            y_scores = []\n\n            for leads_batch, target_batch in tqdm(valid_dataloader):\n                leads_batch = leads_batch.to(device).float()\n                target_batch = target_batch.to(device).float()\n\n                pred_logit_batch = averaged_model(leads_batch)[0]\n                pred_probs_batch = pred_logit_batch.softmax(dim=1)\n\n                pred_probs_batch = torch.clip(pred_probs_batch, min=eps, max=1-eps)\n\n                loss = torch.nn.functional.kl_div(pred_probs_batch.log(), target_batch, reduction='sum')\n                y_scores.append(loss.item())\n\n            score = np.sum(y_scores) / len(valid_dataset)\n\n            if score < best_score:\n                best_score = score\n                best_state_dict = copy.deepcopy(averaged_model.state_dict())\n\n            print('[Epoch {}/{}] [Score: {}] [Best score: {}]'.format(epoch+1, nepochs, score, best_score))\n\naveraged_model.load_state_dict(best_state_dict)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T08:00:35.456347Z","iopub.status.idle":"2024-03-18T08:00:35.456797Z","shell.execute_reply.started":"2024-03-18T08:00:35.456567Z","shell.execute_reply":"2024-03-18T08:00:35.456586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Create submission**","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\nsub = pd.DataFrame(columns=['eeg_id', *classes])\n\nmodel.eval()\n\nwith torch.no_grad():\n    for eidx in test_data:\n        eeg = test_data[eidx]\n        eeg_batch = torch.tensor([eeg.T]).to(device).float()\n\n        logits = model(eeg_batch)[0]\n        probs = torch.round(logits.softmax(dim=1)[0], decimals=17).cpu().numpy()\n\n        sub.loc[len(sub)] = (eidx, *probs)\n\nsub['eeg_id'] = sub['eeg_id'].astype(int)\n\nsub.to_csv('submission.csv', index=None)","metadata":{"execution":{"iopub.status.busy":"2024-03-18T08:00:35.458480Z","iopub.status.idle":"2024-03-18T08:00:35.458907Z","shell.execute_reply.started":"2024-03-18T08:00:35.458686Z","shell.execute_reply":"2024-03-18T08:00:35.458704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2024-03-18T08:00:35.460263Z","iopub.status.idle":"2024-03-18T08:00:35.460770Z","shell.execute_reply.started":"2024-03-18T08:00:35.460505Z","shell.execute_reply":"2024-03-18T08:00:35.460524Z"},"trusted":true},"execution_count":null,"outputs":[]}]}