{"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":11989517,"sourceType":"datasetVersion","datasetId":7541117},{"sourceId":417334,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":339701,"modelId":360813}],"dockerImageVersionId":31011,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\n\nimport random\n\nimport numpy as np\nimport polars as pl\n\nfrom scipy import signal\nfrom scipy.signal import butter, lfilter\nfrom tqdm import tqdm\nfrom glob import glob\n\nfrom dataclasses import dataclass, field\nfrom pathlib import Path\nfrom typing import Dict, List, Sequence, Tuple, Union","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:08.656033Z","iopub.execute_input":"2025-05-29T17:22:08.656258Z","iopub.status.idle":"2025-05-29T17:22:13.571075Z","shell.execute_reply.started":"2025-05-29T17:22:08.656232Z","shell.execute_reply":"2025-05-29T17:22:13.570475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/math156-model-1d\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:13.572721Z","iopub.execute_input":"2025-05-29T17:22:13.573006Z","iopub.status.idle":"2025-05-29T17:22:13.576586Z","shell.execute_reply.started":"2025-05-29T17:22:13.572989Z","shell.execute_reply":"2025-05-29T17:22:13.575864Z"}},"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":{"iopub.status.busy":"2025-05-29T17:22:13.577459Z","iopub.execute_input":"2025-05-29T17:22:13.577967Z","iopub.status.idle":"2025-05-29T17:22:13.595922Z","shell.execute_reply.started":"2025-05-29T17:22:13.577950Z","shell.execute_reply":"2025-05-29T17:22:13.595309Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"#from model import resnet50_1d\n#from model2 import EEGNet\nfrom model3 import resnext","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:13.596495Z","iopub.execute_input":"2025-05-29T17:22:13.596690Z","iopub.status.idle":"2025-05-29T17:22:13.620806Z","shell.execute_reply.started":"2025-05-29T17:22:13.596675Z","shell.execute_reply":"2025-05-29T17:22:13.620316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#model = resnet50_1d(num_classes=6, in_channels=19)\n#model = EEGNet()\nmodel = resnext(num_classes=6, in_channels=19, width_mult = 1.0)\nmodel.load_state_dict(fix_keys(torch.load(\"/kaggle/input/hms-1dmodels-test/pytorch/default/6/best_model3.pth\")))\nmodel.to(\"cuda\")\nprint(\"number of parameters: \", sum(p.numel() for p in model.parameters()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:13.621428Z","iopub.execute_input":"2025-05-29T17:22:13.621664Z","iopub.status.idle":"2025-05-29T17:22:27.937411Z","shell.execute_reply.started":"2025-05-29T17:22:13.621648Z","shell.execute_reply":"2025-05-29T17:22:27.936736Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # basic\n    seed = 42\n    fold = 0\n    \n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    batch_size=32\n    autocast=False # used for training, not for validation\n\n    dataset_path = \"/kaggle/input/hms-harmful-brain-activity-classification\"\n\n    target_cols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:27.938149Z","iopub.execute_input":"2025-05-29T17:22:27.938360Z","iopub.status.idle":"2025-05-29T17:22:27.943129Z","shell.execute_reply.started":"2025-05-29T17:22:27.938343Z","shell.execute_reply":"2025-05-29T17:22:27.942398Z"}},"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 butter_lowpass_filter(data: np.ndarray, cutoff: float, fs: int, order: int = 4) -> np.ndarray:\n    b, a = butter(order, cutoff, fs=fs, btype=\"low\")\n    return lfilter(b, a, data)\n\ndef eeg_from_parquet(\n    parquet_path: str,\n) -> np.ndarray:\n    eeg = pl.read_parquet(parquet_path).drop(\"EKG\").cast(pl.Float32)\n\n    # 2) centre-crop to CFG.nsamples rows (gracefully handles short files)\n    rows = len(eeg)\n    offset = max((rows - 10000) // 2, 0)\n    eeg_slice = eeg[offset : offset + 10000]\n\n    # 3) convert to NumPy and fill NaNs\n    data = eeg_slice.to_numpy()\n    col_mean = np.nanmean(data, axis=0)\n    # columns that are entirely NaN → col_mean becomes NaN, so we replace later\n    nan_rows, nan_cols = np.where(np.isnan(data))\n    data[nan_rows, nan_cols] = col_mean[nan_cols]\n    data = np.nan_to_num(data, nan=0.0)  # all-NaN columns → 0\n\n    return data\n\nset_seed(CFG.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:27.945588Z","iopub.execute_input":"2025-05-29T17:22:27.945804Z","iopub.status.idle":"2025-05-29T17:22:27.977418Z","shell.execute_reply.started":"2025-05-29T17:22:27.945788Z","shell.execute_reply":"2025-05-29T17:22:27.976891Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# load data + split + Preprocessing","metadata":{}},{"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":{"iopub.status.busy":"2025-05-29T17:22:27.978239Z","iopub.execute_input":"2025-05-29T17:22:27.978975Z","iopub.status.idle":"2025-05-29T17:22:28.179889Z","shell.execute_reply.started":"2025-05-29T17:22:27.978955Z","shell.execute_reply":"2025-05-29T17:22:28.179274Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class EEGDataset(Dataset):\n    def __init__(\n        self,\n        df: pl.DataFrame,\n        train: bool = True,\n    ) -> None:\n        self.all_data = df.with_columns(\n                        pl.col(\"path\").map_elements(eeg_from_parquet).alias('eeg')\n                        ).select(['eeg', 'eeg_id']).to_dicts()\n        self.training = train\n\n    def __len__(self) -> int:\n        return len(self.all_data)\n\n    def __getitem__(self, idx: int):\n        row = self.all_data[idx]\n        y = row['eeg_id']\n        X = np.array(row['eeg']) # shape (nsamples, raw_channels)\n            \n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X) / 32.0  # scale down\n        X = butter_lowpass_filter(\n            X, cutoff=20, fs=200, order=6\n        )\n\n        return torch.tensor(X, dtype=torch.float32).permute(1,0),y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:28.180589Z","iopub.execute_input":"2025-05-29T17:22:28.180845Z","iopub.status.idle":"2025-05-29T17:22:28.186661Z","shell.execute_reply.started":"2025-05-29T17:22:28.180820Z","shell.execute_reply":"2025-05-29T17:22:28.186011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset = EEGDataset(df, train=False)\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, drop_last=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:28.187286Z","iopub.execute_input":"2025-05-29T17:22:28.187532Z","iopub.status.idle":"2025-05-29T17:22:28.301751Z","shell.execute_reply.started":"2025-05-29T17:22:28.187508Z","shell.execute_reply":"2025-05-29T17:22:28.300994Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# train, valid (one-epoch)","metadata":{}},{"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":{"iopub.status.busy":"2025-05-29T17:22:28.302509Z","iopub.execute_input":"2025-05-29T17:22:28.302740Z","iopub.status.idle":"2025-05-29T17:22:28.307581Z","shell.execute_reply.started":"2025-05-29T17:22:28.302723Z","shell.execute_reply":"2025-05-29T17:22:28.306846Z"}},"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 eeg, eeg_id in test_loader:\n        eeg   = eeg.to(CFG.device)\n\n        log_pred = model(eeg).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":{"iopub.status.busy":"2025-05-29T17:22:28.308448Z","iopub.execute_input":"2025-05-29T17:22:28.308722Z","iopub.status.idle":"2025-05-29T17:22:29.694402Z","shell.execute_reply.started":"2025-05-29T17:22:28.308697Z","shell.execute_reply":"2025-05-29T17:22:29.693833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submit_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-29T17:22:29.695104Z","iopub.execute_input":"2025-05-29T17:22:29.695286Z","iopub.status.idle":"2025-05-29T17:22:29.704604Z","shell.execute_reply.started":"2025-05-29T17:22:29.695273Z","shell.execute_reply":"2025-05-29T17:22:29.703866Z"}},"outputs":[],"execution_count":null}]}