{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport random\nimport gc\n\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport matplotlib.pyplot as plt\n\nfrom scipy import signal\nfrom tqdm import tqdm\n\nfrom torchaudio import transforms as T\nimport albumentations as A\n\nfrom torchvision.transforms import v2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:55.716079Z","iopub.execute_input":"2025-05-20T20:12:55.716381Z","iopub.status.idle":"2025-05-20T20:12:55.721648Z","shell.execute_reply.started":"2025-05-20T20:12:55.716361Z","shell.execute_reply":"2025-05-20T20:12:55.720834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # basic\n    model_name = \"resnet34\"\n    seed = 42\n    fold = 0\n\n    # training setting\n    n_epoch = 30\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    batch_size=64\n    lr = 1e-5\n    img_size = (257, 600)\n    train_transform=v2.Resize(img_size)\n    valid_transform=v2.Resize(img_size)\n    autocast=True # used for training, not for validation\n\n    dataset_path = \"/kaggle/input/hms-harmful-brain-activity-classification\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:55.722984Z","iopub.execute_input":"2025-05-20T20:12:55.723645Z","iopub.status.idle":"2025-05-20T20:12:55.742405Z","shell.execute_reply.started":"2025-05-20T20:12:55.723620Z","shell.execute_reply":"2025-05-20T20:12:55.741853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''# if TPU\n#!pip install cloud-tpu-client==0.10 torch==1.12.0 https://storage.googleapis.com/tpu-pytorch/wheels/cuda/112/torch_xla-1.12-cp37-cp37m-linux_x86_64.whl --force-reinstall \nimport torch_xla\nimport torch_xla.core.xla_model as xm\n\nclass CFG:\n    # basic\n    model_name = \"resnet34\"\n    seed = 42\n    fold = 0\n\n    # training setting\n    n_epoch = 50\n    device = xm.xla_device()\n    batch_size=64\n    lr = 1e-6\n    img_size = (257, 600)\n    train_transform=v2.Resize(img_size)\n    valid_transform=v2.Resize(img_size)\n    autocast=True # used for training, not for validation\n\n    dataset_path = \"/kaggle/input/hms-harmful-brain-activity-classification\"'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:55.743013Z","iopub.execute_input":"2025-05-20T20:12:55.743236Z","iopub.status.idle":"2025-05-20T20:12:55.758233Z","shell.execute_reply.started":"2025-05-20T20:12:55.743214Z","shell.execute_reply":"2025-05-20T20:12:55.757681Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"def set_seed(seed):\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True\n    \n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n\ndef eeg2spec(eeg):\n    # x: (C,T) or (B, C, T)\n    input_len = eeg.shape[-1]\n    transform = T.Spectrogram(n_fft = 512,\n                          win_length = 64,\n                          hop_length = input_len // 600, # \n                          power = 1)\n    spec = transform(eeg)**0.8\n    spec = torch.nan_to_num(spec)\n    spec = F.normalize(spec)\n    return spec\n\nset_seed(CFG.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:55.758885Z","iopub.execute_input":"2025-05-20T20:12:55.759058Z","iopub.status.idle":"2025-05-20T20:12:55.776896Z","shell.execute_reply.started":"2025-05-20T20:12:55.759043Z","shell.execute_reply":"2025-05-20T20:12:55.776405Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# load data + split + Preprocessing","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedGroupKFold\nsgkf = StratifiedGroupKFold(n_splits=5, shuffle=True, random_state=CFG.seed)\n\n# Load training data\ndf = pl.read_csv(f\"{CFG.dataset_path}/train.csv\")\n\nfor fold, (train_idx, valid_idx) in enumerate(sgkf.split(df, y=df[\"expert_consensus\"], groups=df[\"patient_id\"])):\n    if fold == CFG.fold:\n        break\n\ntrain_df = df[train_idx]\nvalid_df = df[valid_idx]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:55.779094Z","iopub.execute_input":"2025-05-20T20:12:55.779266Z","iopub.status.idle":"2025-05-20T20:12:56.865104Z","shell.execute_reply.started":"2025-05-20T20:12:55.779252Z","shell.execute_reply":"2025-05-20T20:12:56.864534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"labels = [\"seizure\", \"lpd\", \"gpd\", \"lrda\", \"grda\", \"other\"]\n\ntrain_labels = train_df.group_by('eeg_id', maintain_order=True).agg([\n    *[pl.col(lbl+'_vote').sum() for lbl in labels],\n    pl.len().alias('total_vote') \n]).with_columns(\n    pl.sum_horizontal([pl.col(lbl+'_vote') for lbl in labels]).alias(\"total_vote\").cast(pl.Float64)\n).with_columns(\n    *[pl.col(lbl+'_vote') / pl.col('total_vote') for lbl in labels],\n    pl.concat_str(\n    [\n            pl.lit(f\"{CFG.dataset_path}/train_eegs/\"),\n            pl.col('eeg_id').cast(pl.String),\n            pl.lit(\".parquet\"),\n        ],\n    ).alias(\"path\"),\n).drop(\"total_vote\")\n\nvalid_labels = valid_df.group_by('eeg_id', maintain_order=True).agg([\n    *[pl.col(lbl+'_vote').sum() for lbl in labels],\n    pl.len().alias('total_vote') \n]).with_columns(\n    pl.sum_horizontal([pl.col(lbl+'_vote') for lbl in labels]).alias(\"total_vote\").cast(pl.Float64)\n).with_columns(\n    *[pl.col(lbl+'_vote') / pl.col('total_vote') for lbl in labels],\n    pl.concat_str(\n    [\n            pl.lit(f\"{CFG.dataset_path}/train_eegs/\"),\n            pl.col('eeg_id').cast(pl.String),\n            pl.lit(\".parquet\"),\n        ],\n    ).alias(\"path\"),\n).drop(\"total_vote\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:56.865725Z","iopub.execute_input":"2025-05-20T20:12:56.865923Z","iopub.status.idle":"2025-05-20T20:12:56.948898Z","shell.execute_reply.started":"2025-05-20T20:12:56.865907Z","shell.execute_reply":"2025-05-20T20:12:56.948378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del df, train_df, valid_df\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:56.949408Z","iopub.execute_input":"2025-05-20T20:12:56.949620Z","iopub.status.idle":"2025-05-20T20:12:57.154106Z","shell.execute_reply.started":"2025-05-20T20:12:56.949603Z","shell.execute_reply":"2025-05-20T20:12:57.153190Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Show created Spectrogram","metadata":{}},{"cell_type":"code","source":"np.random.seed(40)\nfor _ in np.random.randint(len(train_labels), size=(10)):\n    idx = int(_)\n    data = pl.read_parquet(train_labels['path'][idx])\n    spec = eeg2spec(data.drop(\"EKG\").to_torch().transpose(1,0))\n    print(f\"seizure_vote: {train_labels['seizure_vote'][idx]:.2f} | lpd_vote: {train_labels['lpd_vote'][idx]:.2f} | gpd_vote: {train_labels['gpd_vote'][idx]:.2f} | lrda_vote: {train_labels['lrda_vote'][idx]:.2f} | grda_vote: {train_labels['grda_vote'][idx]:.2f} | other_vote: {train_labels['other_vote'][idx]:.2f}\")\n    plt.imshow(torch.cat([spec[i][:40, :] for i in range(len(spec))]))\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:57.154998Z","iopub.execute_input":"2025-05-20T20:12:57.155333Z","iopub.status.idle":"2025-05-20T20:12:59.959378Z","shell.execute_reply.started":"2025-05-20T20:12:57.155312Z","shell.execute_reply":"2025-05-20T20:12:59.958382Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class HmsDataset(Dataset):\n    def __init__(self, labels_df: pl.DataFrame, train=False, transform=None):\n        '''\n        in train/valid - set train True\n        in inference - set train False\n        '''\n        self.labels_df = labels_df\n        self.paths = labels_df['path'].to_list()\n        self.train = train\n        if train:\n            self.targets = labels_df.select([\n                    'seizure_vote', 'lpd_vote', 'gpd_vote',\n                    'lrda_vote', 'grda_vote', 'other_vote'\n                ]).to_torch()\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx: int):\n        df = pl.read_parquet(self.paths[idx])\n        eeg = df.drop(\"EKG\").to_torch().transpose(1, 0)\n        spec = eeg2spec(eeg)\n        if self.transform is not None:\n            spec = self.transform(spec)\n        if self.train:\n            target = self.targets[idx]\n            return spec, target\n        else:\n            return spec","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:59.960424Z","iopub.execute_input":"2025-05-20T20:12:59.960805Z","iopub.status.idle":"2025-05-20T20:12:59.969087Z","shell.execute_reply.started":"2025-05-20T20:12:59.960775Z","shell.execute_reply":"2025-05-20T20:12:59.968615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = HmsDataset(train_labels, train=True, transform=CFG.train_transform)\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True)\nvalid_dataset = HmsDataset(valid_labels, train=True, transform=CFG.valid_transform)\nvalid_loader = DataLoader(valid_dataset, batch_size=CFG.batch_size, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:59.970116Z","iopub.execute_input":"2025-05-20T20:12:59.970459Z","iopub.status.idle":"2025-05-20T20:12:59.995639Z","shell.execute_reply.started":"2025-05-20T20:12:59.970430Z","shell.execute_reply":"2025-05-20T20:12:59.994736Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"import timm\nmodel = timm.create_model(CFG.model_name, pretrained=True, num_classes=6, in_chans=19).to(CFG.device)\nprint(\"number of parameters: \", sum(p.numel() for p in model.parameters()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:12:59.996866Z","iopub.execute_input":"2025-05-20T20:12:59.997128Z","iopub.status.idle":"2025-05-20T20:13:03.957265Z","shell.execute_reply.started":"2025-05-20T20:12:59.997104Z","shell.execute_reply":"2025-05-20T20:13:03.956600Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Optimizer","metadata":{}},{"cell_type":"code","source":"optimizer = optim.SGD(model.parameters(), lr=CFG.lr)\nkl_loss = nn.KLDivLoss(reduction=\"batchmean\")\n# use scheduler here? - skip this time","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:13:03.957919Z","iopub.execute_input":"2025-05-20T20:13:03.958112Z","iopub.status.idle":"2025-05-20T20:13:03.962636Z","shell.execute_reply.started":"2025-05-20T20:13:03.958096Z","shell.execute_reply":"2025-05-20T20:13:03.961899Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# train, valid (one-epoch)","metadata":{}},{"cell_type":"code","source":"def train(model, train_loader, optimizer, criterion, epoch):\n    model.train()\n    running_loss = 0.0\n\n    for spec, target in train_loader:\n        spec   = spec.to(CFG.device)\n        target = target.to(CFG.device)\n\n        optimizer.zero_grad()\n        with torch.autocast(device_type=\"cuda\", enabled=CFG.autocast):\n            log_pred = model(spec).log_softmax(dim=1)\n            target_clipped = torch.clamp(target, 1e-8, 1.0)\n            loss = criterion(log_pred, target_clipped)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n\n    avg_loss = running_loss / len(train_loader)\n    print(f\"*** Epoch {epoch} TRAINING COMPLETE. Avg Loss: {avg_loss:.4f} ***\")\n    torch.cuda.empty_cache(); gc.collect()\n    return avg_loss\n\n\ndef valid(model, valid_loader, criterion, epoch):\n    model.eval()\n    all_log_pred, all_target = [], []\n\n    with torch.no_grad(), torch.autocast(device_type=\"cuda\", enabled=False):\n        for spec, target in valid_loader:\n            spec   = spec.to(CFG.device)\n            target = target.to(CFG.device)\n\n            log_pred = model(spec).log_softmax(dim=1)\n            all_log_pred.append(log_pred)\n            all_target.append(target)\n\n        log_pred = torch.cat(all_log_pred, dim=0)\n        target   = torch.cat(all_target , dim=0)\n        target   = torch.clamp(target, 1e-8, 1.0)\n\n        val_loss = criterion(log_pred, target)\n\n    print(f\"=== Epoch {epoch} VALIDATION: KL Divergence = {val_loss:.4f} ===\")\n    torch.cuda.empty_cache(); gc.collect()\n    return val_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:13:03.963419Z","iopub.execute_input":"2025-05-20T20:13:03.963810Z","iopub.status.idle":"2025-05-20T20:13:03.994541Z","shell.execute_reply.started":"2025-05-20T20:13:03.963785Z","shell.execute_reply":"2025-05-20T20:13:03.993878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = kl_loss\nepoch = -1\n\nmodel.train()\nrunning_loss = 0.0\n\nfor spec, target in train_loader:\n    spec   = spec.to(CFG.device)\n    target = target.to(CFG.device)\n\n    optimizer.zero_grad()\n    with torch.autocast(device_type=\"cuda\", enabled=CFG.autocast):\n        log_pred = model(spec).log_softmax(dim=1)\n        target_clipped = torch.clamp(target, 1e-8, 1.0)\n        loss = criterion(log_pred, target_clipped)\n\n    loss.backward()\n    optimizer.step()\n\n    running_loss += loss.item()\n\navg_loss = running_loss / len(train_loader)\nprint(f\"=== Epoch {epoch} TRAINING COMPLETE. Avg Loss: {avg_loss:.4f} ===\")\ntorch.cuda.empty_cache(); gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:13:03.996913Z","iopub.execute_input":"2025-05-20T20:13:03.997181Z","iopub.status.idle":"2025-05-20T20:31:25.383830Z","shell.execute_reply.started":"2025-05-20T20:13:03.997164Z","shell.execute_reply":"2025-05-20T20:31:25.383278Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_loss = np.inf\n\nfor epoch in range(1, CFG.n_epoch+1):\n    avg_loss = train(model, train_loader, optimizer, kl_loss, epoch)\n    avg_val_loss = valid(model, valid_loader, kl_loss, epoch)\n    if avg_val_loss < best_loss:\n        best_loss = avg_val_loss\n        torch.save(model.state_dict(), \"best_model.pth\")\n        print(f\">>> New best model saved (KL={best_loss:.4f}) <<<\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-20T20:31:25.384541Z","iopub.execute_input":"2025-05-20T20:31:25.384888Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# TODO\n- pretrained resnet 18, 34, 50 - use timm\n- not pretrained resnet 18,34,50 - use timm\n- create submission (inference on test data) pipeline (let me know if you are not sure how to submit)\n## What to Report\nparameter size, KL-Div (valid & test)","metadata":{}}]}