{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":418628,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":341451,"modelId":362723}],"dockerImageVersionId":31041,"isInternetEnabled":false,"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-30T18:02:34.404648Z","iopub.execute_input":"2025-05-30T18:02:34.405205Z","iopub.status.idle":"2025-05-30T18:02:34.409726Z","shell.execute_reply.started":"2025-05-30T18:02:34.405181Z","shell.execute_reply":"2025-05-30T18:02:34.408978Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nfrom pathlib import Path\n# update for 2d cnn\n# sys.path.append(\"/kaggle/input/math156-model-1d\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-30T18:02:34.410962Z","iopub.execute_input":"2025-05-30T18:02:34.411288Z","iopub.status.idle":"2025-05-30T18:02:34.427206Z","shell.execute_reply.started":"2025-05-30T18:02:34.411264Z","shell.execute_reply":"2025-05-30T18:02:34.426586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fix_keys(loaded_dict):\n    return {k.replace(\"_orig_mod.\", \"\"): v for k, v in loaded_dict.items()}","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    # basic\n    model_name = \"resnet18\"\n    seed = 42\n    fold = 0\n\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    batch_size=64\n    img_size = (257, 600)\n    train_transform=v2.Resize(img_size)\n    valid_transform=v2.Resize(img_size)\n    autocast=False # used for training, not for validation\n\n    dataset_path = \"/kaggle/input/hms-harmful-brain-activity-classification\"\n    target_cols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# update for 2d cnn\n\nimport timm\n\nmodel = timm.create_model(CFG.model_name, pretrained=False, num_classes=6, in_chans=19).to(CFG.device)\n# modify below line for 2d \nmodel.load_state_dict(fix_keys(torch.load(f\"/kaggle/input/resnet_spec/pytorch/default/1/{CFG.model_name}_best_model.pth\")))\nmodel.to(CFG.device)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.646Z"}},"outputs":[],"execution_count":null},{"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":{"execution_failed":"2025-05-30T18:02:58.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pl.read_csv(f\"{CFG.dataset_path}/test.csv\").with_columns(pl.concat_str(\n    [\n            pl.lit(f\"{CFG.dataset_path}/test_eegs/\"),\n            pl.col('eeg_id').cast(pl.String),\n            pl.lit(\".parquet\"),\n        ],\n    ).alias(\"path\")\n           )","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.647Z"}},"outputs":[],"execution_count":null},{"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        \n        self.targets = labels_df.select([\n                'eeg_id'\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        target = self.targets[idx]\n        return spec, target","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset = HmsDataset(df, train=False, transform=CFG.valid_transform)\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, drop_last=False)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def _fix_row(row: np.ndarray) -> np.ndarray:\n    \"\"\"Return a copy whose elements sum to 1.0 by tweaking one entry.\"\"\"\n    s = row.sum(dtype=np.float64)\n    delta = 1.0 - s\n    if abs(delta) < 1e-6:           # already OK\n        return row\n\n    # choose index to adjust\n    idx = row.argmin() if delta > 0 else row.argmax()\n    row[idx] += delta               # add (or subtract) the difference\n    # numerical safety - clip into [0,1]\n    row[idx] = np.clip(row[idx], 0.0, 1.0)\n    return row","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nall_log_pred, all_eeg_id = [], []\n\nwith torch.no_grad(), torch.autocast(device_type=\"cuda\", enabled=False):\n    for spec, eeg_id in test_loader:\n        spec = spec.to(CFG.device)\n\n        log_pred = model(spec).log_softmax(dim=1).cpu()\n        all_log_pred.append(log_pred)\n        all_eeg_id.append(eeg_id)\n\nlog_preds = torch.cat(all_log_pred, dim=0)          # (N, n_classes)\neeg_ids   = torch.cat(all_eeg_id , dim=0)           # (N,)\n\nlog_preds = log_preds.reshape(-1, log_preds.size(-1))  # (N, C)\neeg_ids   = eeg_ids.reshape(-1)                        # (N,)\n\nprobs = log_preds.exp()\n\n\nprobs_np = probs.cpu().numpy().astype(np.float32)         # shape = (N, C)\nprobs_np = np.apply_along_axis(_fix_row, 1, probs_np)     # normalise rows\n\n\nsubmit_df = pl.DataFrame({\n    \"eeg_id\":   eeg_ids.cpu().numpy().astype(\"int64\"),\n    **{c: probs_np[:, i] for i, c in enumerate(CFG.target_cols)}\n})\n\nsubmit_df.write_csv(\"submission.csv\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submit_df","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-30T18:02:58.647Z"}},"outputs":[],"execution_count":null}]}