{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"},"colab":{"gpuType":"T4","provenance":[]},"accelerator":"GPU","papermill":{"default_parameters":{},"duration":33298.215016,"end_time":"2026-08-23T00:03:25.618129+00:00","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-08-22T14:48:27.403113+00:00","version":"2.7.0"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"13ca3fa0474f41cf87ee345330ea361a":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HTMLView","description":"","description_allow_html":false,"layout":"IPY_MODEL_7e886d1c21154437b611af78b5c79484","placeholder":"​","style":"IPY_MODEL_57fe29a50adc4e9e8ae941a1b4a73e04","tabbable":null,"tooltip":null,"value":" 21.4M/21.4M [00:01&lt;00:00, 15.4MB/s]"}},"23b94bb2584c44a7a900e97dd5d90aad":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"532a3268352e401ea8f5504491cbd64c":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HTMLView","description":"","description_allow_html":false,"layout":"IPY_MODEL_dbba9ff744d54d4c8cff9e117937402d","placeholder":"​","style":"IPY_MODEL_dbbb9f9b29ea413e9f61e812ff61f74b","tabbable":null,"tooltip":null,"value":"model.safetensors: 100%"}},"575ccf7abf914af9bbb534f6543d6329":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_532a3268352e401ea8f5504491cbd64c","IPY_MODEL_d6116813a7f84dc8ad3dd01a8b62b7be","IPY_MODEL_13ca3fa0474f41cf87ee345330ea361a"],"layout":"IPY_MODEL_e193d4d02dd649139f2a3d80ba7c5965","tabbable":null,"tooltip":null}},"57fe29a50adc4e9e8ae941a1b4a73e04":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","background":null,"description_width":"","font_size":null,"text_color":null}},"7e886d1c21154437b611af78b5c79484":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"b6791662ff8744e7b4d864198d351c0f":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"d6116813a7f84dc8ad3dd01a8b62b7be":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"ProgressView","bar_style":"success","description":"","description_allow_html":false,"layout":"IPY_MODEL_b6791662ff8744e7b4d864198d351c0f","max":21355344,"min":0,"orientation":"horizontal","style":"IPY_MODEL_23b94bb2584c44a7a900e97dd5d90aad","tabbable":null,"tooltip":null,"value":21355344}},"dbba9ff744d54d4c8cff9e117937402d":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"dbbb9f9b29ea413e9f61e812ff61f74b":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","background":null,"description_width":"","font_size":null,"text_color":null}},"e193d4d02dd649139f2a3d80ba7c5965":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}}},"version_major":2,"version_minor":0}},"description":"Controlled stronger-dropout ablation on the winning sqrt-sampler pipeline. Baseline dropout was 0.30; this version uses 0.50 and leaves sampler/data/loss/schedule unchanged.","experiment":"Experiment: sqrt-inverse-frequency sampler + dropout 0.50 (1 fold)","kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":59093},{"sourceType":"datasetVersion","sourceId":18757649}],"dockerImageVersionId":31400,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"path = \"/kaggle/input/competitions/hms-harmful-brain-activity-classification/train.csv\"","metadata":{"execution":{"iopub.status.busy":"2026-09-11T06:03:08.026134Z","iopub.execute_input":"2026-09-11T06:03:08.026417Z","iopub.status.idle":"2026-09-11T06:03:08.033490Z","shell.execute_reply.started":"2026-09-11T06:03:08.026394Z","shell.execute_reply":"2026-09-11T06:03:08.032749Z"},"id":"wv7CSuJl1Ldb","papermill":{"duration":0.013985,"end_time":"2026-08-22T14:48:29.77207+00:00","exception":false,"start_time":"2026-08-22T14:48:29.758085+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# No extra loss package is required for the main experiment.\n# The paper's primary objective is soft-target KL divergence with a standard\n# softmax output, so the auxiliary entmax loss is disabled.\nprint(\"Using softmax + KL divergence only for the main modality comparison.\")\n","metadata":{"execution":{"iopub.status.busy":"2026-09-11T06:03:08.035331Z","iopub.execute_input":"2026-09-11T06:03:08.035623Z","iopub.status.idle":"2026-09-11T06:03:08.046591Z","shell.execute_reply.started":"2026-09-11T06:03:08.035602Z","shell.execute_reply":"2026-09-11T06:03:08.045799Z"},"papermill":{"duration":5.596102,"end_time":"2026-08-22T14:48:35.372725+00:00","exception":false,"start_time":"2026-08-22T14:48:29.776623+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\ndf = pd.read_csv(path)\ndf.head(10)","metadata":{"execution":{"iopub.status.busy":"2026-09-11T06:03:08.047475Z","iopub.execute_input":"2026-09-11T06:03:08.047786Z","iopub.status.idle":"2026-09-11T06:03:09.432039Z","shell.execute_reply.started":"2026-09-11T06:03:08.047742Z","shell.execute_reply":"2026-09-11T06:03:09.430999Z"},"id":"sDm6ifDl1Lff","outputId":"017960af-a70b-4c11-aa9d-7bf23a3b4b1e","papermill":{"duration":0.989002,"end_time":"2026-08-22T14:48:36.36738+00:00","exception":false,"start_time":"2026-08-22T14:48:35.378378+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"vote_cols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n\n# FIX: the previous version grouped by eeg_id and averaged offset_min/offset_max\n# across ALL labeled sub-windows belonging to that eeg_id, then summed votes\n# across them too. But different sub_ids of the same eeg_id can be separate\n# 50-second labeled windows at different points in the recording - sometimes\n# with different expert_consensus labels entirely. Collapsing them into one\n# row blurred the crop center away from any window an expert actually rated,\n# and mixed votes from unrelated windows into one soft label.\n#\n# Fix: keep one row PER LABELED WINDOW (what train.csv actually encodes),\n# instead of collapsing to one row per eeg_id.\ndf['total_votes'] = df[vote_cols].sum(axis=1)\n\n# FILTER NOISE: Keep windows with at least 8 expert votes for cleaner training targets\nagg_df = df[df['total_votes'] >= 8].copy()\n\n# Keep the original annotator-vote proportions as the TRAINING target.\n# This is the no-smoothing ablation: VOTE_SMOOTH_EPS = 0.0. The same raw\n# proportions are also retained separately for evaluation, so training and\n# raw-vote KL now use the same target definition.\nVOTE_SMOOTH_EPS = 0.0\nn_classes = len(vote_cols)\nagg_df[vote_cols] = (\n    agg_df[vote_cols] / np.maximum(agg_df['total_votes'].to_numpy()[:, None], 1.0)\n)\n\n# sort by eeg_id/spectrogram_id (not required for correctness) purely so the\n# caching loop below reads each parquet file's rows contiguously instead of\n# jumping around, which cuts redundant parquet reads when a file has\n# multiple labeled windows\nagg_df = agg_df.sort_values(['eeg_id', 'spectrogram_id']).reset_index(drop=True)\n\n# Apply the exact same row ordering to the raw targets so they align with the\n# cached samples, whose labels were built from this sorted agg_df.\n_sort_order = df[df['total_votes'] >= 8].sort_values(\n    ['eeg_id', 'spectrogram_id']\n).index.to_numpy()\nraw_vote_targets = (\n    df.loc[_sort_order, vote_cols].to_numpy(dtype=np.float32)\n    / np.maximum(\n        df.loc[_sort_order, vote_cols].to_numpy(dtype=np.float32).sum(axis=1, keepdims=True),\n        1.0\n    )\n)\n\nprint(f\"Raw-vote evaluation targets prepared: {raw_vote_targets.shape}\")\nprint(f\"Total clean labeled windows (total_votes >= 8): {len(agg_df)}\")\nprint(f\"  ({df['eeg_id'].nunique()} unique eeg_ids - previously collapsed to one row each; \"\n      f\"now {len(agg_df)} true per-window rows)\")\n","metadata":{"execution":{"iopub.status.busy":"2026-09-11T06:03:09.433369Z","iopub.execute_input":"2026-09-11T06:03:09.433836Z","iopub.status.idle":"2026-09-11T06:03:09.519478Z","shell.execute_reply.started":"2026-09-11T06:03:09.433787Z","shell.execute_reply":"2026-09-11T06:03:09.518487Z"},"id":"pXfIlSA18fr6","outputId":"ad40bfea-514d-464b-8900-949cdd774061","papermill":{"duration":0.07413,"end_time":"2026-08-22T14:48:36.447217+00:00","exception":false,"start_time":"2026-08-22T14:48:36.373087+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# No-smoothing target sanity check\n_train_target = agg_df[vote_cols].to_numpy(dtype=np.float32)\nassert np.allclose(_train_target.sum(axis=1), 1.0, atol=1e-5)\nassert np.all(_train_target >= 0.0)\nprint(f'No-smoothing training targets: {_train_target.shape}')\nprint(f'Rows with at least one zero target: {(~(_train_target > 0)).any(axis=1).mean()*100:.1f}%')\nprint('Training/checkpoint KL and raw-vote evaluation KL now use the same raw vote proportions.')\n","metadata":{"execution":{"iopub.status.busy":"2026-09-11T06:03:09.520718Z","iopub.execute_input":"2026-09-11T06:03:09.521090Z","iopub.status.idle":"2026-09-11T06:03:09.531369Z","shell.execute_reply.started":"2026-09-11T06:03:09.521064Z","shell.execute_reply":"2026-09-11T06:03:09.530184Z"},"papermill":{"duration":0.015018,"end_time":"2026-08-22T14:48:36.467748+00:00","exception":false,"start_time":"2026-08-22T14:48:36.45273+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\nimport numpy as np\nimport cv2\nimport os\n\ncv2.setNumThreads(0)\n\nspec_path = \"/kaggle/input/competitions/hms-harmful-brain-activity-classification/train_spectrograms\"\neeg_path = \"/kaggle/input/competitions/hms-harmful-brain-activity-classification/train_eegs\"\n\nREGIONS = ['LL', 'RL', 'LP', 'RP']\n\nCROP_W = 256                      # pixels fed to the network (represents a 600s window, same as before)\nSECONDS_PER_PX = 600 / CROP_W     # time resolution implied by the original 600s/256px mapping\nJITTER_PX = 32                    # extra pixels cached on each side, giving room for random time-crop at train time\nWIDE_W = CROP_W + 2 * JITTER_PX   # 320px - what actually gets cached to disk\nWIDE_SECONDS = WIDE_W * SECONDS_PER_PX  # ~750s window cached, cropped down to 600s (256px) per sample\nSPEC_TARGET_H = 384   # <- single source of truth for spectrogram row count.\n                       # Change ONLY this number to test a new resolution -\n                       # build_spectrogram_image's default and cell 9's cache\n                       # preallocation/validity-check both read from it.\n\n\ndef build_spectrogram_image(spec_raw, offset, target_h=SPEC_TARGET_H, target_w=WIDE_W):\n    \"\"\"Crop a window around the labeled region - WIDER than the original 600s\n    (target_w=WIDE_W by default) so random_time_crop() has room to jitter at\n    train time. Keeps the 4 brain regions distinct (not blended into one\n    undifferentiated frequency axis), log-transforms + normalizes each region\n    individually, then stacks them into one mosaic image. This still only\n    calls cv2.resize once per sample, at cache-build time - not during\n    training.\"\"\"\n    window_seconds = target_w * SECONDS_PER_PX\n    pad = (window_seconds - 600) / 2\n    start = offset - pad\n    end = offset + 600 + pad\n\n    window = spec_raw.loc[(spec_raw['time'] >= start) & (spec_raw['time'] < end)]\n    if len(window) == 0:\n        window = spec_raw  # fallback if offset doesn't align with any rows\n\n    region_h = target_h // len(REGIONS)\n    region_imgs = []\n    for region in REGIONS:\n        cols = [c for c in spec_raw.columns if c.startswith(region + '_')]\n        img = window[cols].values.T.astype(np.float32)\n        img = np.nan_to_num(img, nan=0.0)\n        img = np.clip(img, np.exp(-4), np.exp(8))\n        img = np.log(img)\n        m, s = img.mean(), img.std()\n        img = (img - m) / (s + 1e-6)\n        img = cv2.resize(img, (target_w, region_h), interpolation=cv2.INTER_AREA)\n        region_imgs.append(img)\n\n    mosaic = np.concatenate(region_imgs, axis=0)\n    return mosaic.astype(np.float32)\n\n\ndef random_time_crop(img, crop_w=CROP_W):\n    \"\"\"Random horizontal crop from the cached WIDE_W-wide mosaic down to\n    crop_w - this is the time-jitter augmentation. Pure array slicing, no\n    resize, so it's cheap enough to run in __getitem__.\"\"\"\n    max_start = img.shape[1] - crop_w\n    start = np.random.randint(0, max_start + 1) if max_start > 0 else 0\n    return img[:, start:start + crop_w]\n\n\ndef center_time_crop(img, crop_w=CROP_W):\n    \"\"\"Deterministic center crop - used for val/test so evaluation is\n    reproducible run to run.\"\"\"\n    start = (img.shape[1] - crop_w) // 2\n    return img[:, start:start + crop_w]\n\n\ndef spec_augment(img, freq_mask_frac=0.15, time_mask_frac=0.15, n_freq_masks=1, n_time_masks=2):\n    \"\"\"Standard SpecAugment: zero out random frequency bands and time slices.\n    Applied post-crop, on the final 256 x CROP_W image the network actually\n    sees, so mask sizes stay meaningful regardless of crop_w.\"\"\"\n    img = img.copy()\n    h, w = img.shape\n    for _ in range(n_freq_masks):\n        f = int(h * freq_mask_frac * np.random.rand())\n        if f == 0:\n            continue\n        f0 = np.random.randint(0, h - f + 1)\n        img[f0:f0 + f, :] = 0.0\n    for _ in range(n_time_masks):\n        t = int(w * time_mask_frac * np.random.rand())\n        if t == 0:\n            continue\n        t0 = np.random.randint(0, w - t + 1)\n        img[:, t0:t0 + t] = 0.0\n    return img\nFLIP_PROB = 0.5\n\ndef flip_spectrogram_lr(img):\n    \"\"\"Swaps the LL<->RL and LP<->RP region blocks in the spectrogram mosaic\n    (rows, since REGIONS = ['LL','RL','LP','RP'] are stacked vertically).\n    LPD/LRDA abnormalities are lateralized - they show up predominantly in\n    one hemisphere's channels. Mirroring L<->R is a label-preserving\n    augmentation (the class doesn't depend on which side it's on) that\n    forces the model to learn 'abnormality on one side' as the pattern,\n    instead of overfitting to a left/right split in the training data.\"\"\"\n    h = img.shape[0]\n    region_h = h // 4\n    flipped = img.copy()\n    flipped[0:region_h] = img[region_h:2 * region_h]\n    flipped[region_h:2 * region_h] = img[0:region_h]\n    flipped[2 * region_h:3 * region_h] = img[3 * region_h:4 * region_h]\n    flipped[3 * region_h:4 * region_h] = img[2 * region_h:3 * region_h]\n    return flipped\n\n\ndef flip_eeg_lr(sig):\n    \"\"\"Same swap for the EEG montage - BIPOLAR_CHAINS is ordered\n    LL, RL, LP, RP (4 channels each), so this mirrors flip_spectrogram_lr\n    on the channel axis instead of the frequency axis. Must be applied\n    together with flip_spectrogram_lr on the same sample - the two\n    branches represent the same physical hemispheres, so flipping only one\n    would feed the fused model contradictory left/right information.\"\"\"\n    c = sig.shape[0]\n    cpr = c // 4\n    flipped = sig.copy()\n    flipped[0:cpr] = sig[cpr:2 * cpr]\n    flipped[cpr:2 * cpr] = sig[0:cpr]\n    flipped[2 * cpr:3 * cpr] = sig[3 * cpr:4 * cpr]\n    flipped[3 * cpr:4 * cpr] = sig[2 * cpr:3 * cpr]\n    return flipped\n\n# ============================================================================\n# EEG branch: raw signal, mirrors the experts' actual protocol.\n#\n# Per the competition's data description, each labeled row corresponds to\n# experts viewing a fixed 50-second EEG window and rating the event type for\n# the MIDDLE 10 SECONDS of it (the spectrogram is 10-min/600s context around\n# the same point - already handled above). This section builds that raw-EEG\n# 50s -> center-10s pathway as a second model input, fused with the\n# spectrogram branch at training time.\n# ============================================================================\n\nEEG_SR = 200                # Hz - EEGs are resampled to 200Hz per the competition docs\nEEG_FULL_SEC = 50           # exact window length experts were shown\nEEG_TARGET_SEC = 10         # experts' actually-rated region: the middle 10s\nEEG_TARGET_LEN = EEG_TARGET_SEC * EEG_SR             # 2000 samples fed to the network\nEEG_JITTER_SEC = 2.0        # +/- jitter room for train-time augmentation only\nEEG_JITTER_LEN = int(EEG_JITTER_SEC * EEG_SR)        # 400 samples\nEEG_WIDE_LEN = EEG_TARGET_LEN + 2 * EEG_JITTER_LEN   # 2800 samples (~14s) cached to disk\n\n# Standard \"double banana\" bipolar montage: 4 chains of 5 electrodes each,\n# differenced consecutively -> 4 diffs/chain x 4 chains = 16 channels.\nBIPOLAR_CHAINS = {\n    'LL': ['Fp1', 'F7', 'T3', 'T5', 'O1'],\n    'RL': ['Fp2', 'F8', 'T4', 'T6', 'O2'],\n    'LP': ['Fp1', 'F3', 'C3', 'P3', 'O1'],\n    'RP': ['Fp2', 'F4', 'C4', 'P4', 'O2'],\n}\n\n\ndef build_eeg_window(eeg_raw, offset_seconds, sr=EEG_SR, full_sec=EEG_FULL_SEC, wide_len=EEG_WIDE_LEN):\n    \"\"\"Mirrors build_spectrogram_image: the expert only ever saw a fixed\n    50-second EEG window (offset_seconds -> offset_seconds + 50) and rated\n    the middle 10s of it. We cut a WIDE_LEN-sample segment centered on that\n    same midpoint (well inside the 50s bound) so random_time_crop_1d has\n    jitter room at train time; center_time_crop_1d recovers the exact\n    middle-10s segment the expert actually rated for eval.\"\"\"\n    start_sample = int(round(offset_seconds * sr))\n    mid_sample = start_sample + int(full_sec * sr) // 2\n    half_wide = wide_len // 2\n    lo = max(0, mid_sample - half_wide)\n    hi = lo + wide_len\n\n    channels = []\n    for chain in BIPOLAR_CHAINS.values():\n        for a, b in zip(chain[:-1], chain[1:]):\n            sig = (eeg_raw[a].values - eeg_raw[b].values).astype(np.float32)\n            channels.append(sig)\n    montage = np.stack(channels, axis=0)   # (16, n_samples_in_file)\n\n    if hi > montage.shape[1]:\n        pad = hi - montage.shape[1]\n        montage = np.pad(montage, ((0, 0), (0, pad)), mode='edge')\n    window = montage[:, lo:hi]\n    if window.shape[1] < wide_len:\n        pad = wide_len - window.shape[1]\n        window = np.pad(window, ((0, 0), (0, pad)), mode='edge')\n\n    window = np.nan_to_num(window, nan=0.0)\n    m = window.mean(axis=1, keepdims=True)\n    s = window.std(axis=1, keepdims=True)\n    window = (window - m) / (s + 1e-6)\n    window = np.clip(window, -20, 20)   # guard against rare spike artifacts\n    return window.astype(np.float32)\n\n\ndef random_time_crop_1d(sig, crop_len=EEG_TARGET_LEN):\n    \"\"\"Time-jitter augmentation for train, analogous to random_time_crop above.\"\"\"\n    max_start = sig.shape[1] - crop_len\n    start = np.random.randint(0, max_start + 1) if max_start > 0 else 0\n    return sig[:, start:start + crop_len]\n\n\ndef center_time_crop_1d(sig, crop_len=EEG_TARGET_LEN):\n    \"\"\"Deterministic center crop for val/test - this IS the exact middle-10s\n    window the experts rated.\"\"\"\n    start = (sig.shape[1] - crop_len) // 2\n    return sig[:, start:start + crop_len]\n\n\nclass HMSFusionDataset(Dataset):\n    \"\"\"Returns (spectrogram_image, eeg_signal, label). The spectrogram gets\n    the existing 2D time-crop + SpecAugment; the EEG gets a matching 1D\n    time-crop - random jitter at train time, exact center-10s at eval time\n    (i.e. the same window the experts actually rated).\"\"\"\n    def __init__(self, specs, eeg, labels, indices, train=False, crop_w=CROP_W, eeg_crop_len=EEG_TARGET_LEN):\n        self.specs = specs\n        self.eeg = eeg\n        self.labels = labels\n        self.indices = indices\n        self.train = train\n        self.crop_w = crop_w\n        self.eeg_crop_len = eeg_crop_len\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, i):\n        idx = self.indices[i]\n        img = self.specs[idx]\n        sig = self.eeg[idx]\n\n        if self.train:\n            img = random_time_crop(img, self.crop_w)\n            img = spec_augment(img)\n            sig = random_time_crop_1d(sig, self.eeg_crop_len)\n            if np.random.rand() < FLIP_PROB:\n                img = flip_spectrogram_lr(img)\n                sig = flip_eeg_lr(sig)\n        else:\n            img = center_time_crop(img, self.crop_w)\n            sig = center_time_crop_1d(sig, self.eeg_crop_len)\n\n        img = np.stack([img, img, img])\n        return (torch.tensor(img, dtype=torch.float32),\n                torch.tensor(sig, dtype=torch.float32),\n                torch.tensor(self.labels[idx], dtype=torch.float32))\n","metadata":{"execution":{"iopub.status.busy":"2026-09-11T06:03:09.533694Z","iopub.execute_input":"2026-09-11T06:03:09.534060Z","iopub.status.idle":"2026-09-11T06:03:14.861294Z","shell.execute_reply.started":"2026-09-11T06:03:09.534007Z","shell.execute_reply":"2026-09-11T06:03:14.860552Z"},"id":"ag6raeeW-Sa0","papermill":{"duration":4.303377,"end_time":"2026-08-22T14:48:40.776416+00:00","exception":false,"start_time":"2026-08-22T14:48:36.473039+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Cache compatibility note\nThe attached v4 cache is automatically detected, but it is not silently used here because its CWT representation is incompatible with this notebook's original 0.27 raw-EEG fusion branch. The Experiment 1 result must remain a clean baseline comparison.\n","metadata":{"papermill":{"duration":0.005464,"end_time":"2026-08-22T14:48:40.787437+00:00","exception":false,"start_time":"2026-08-22T14:48:40.781973+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ============================================================================\n# CACHE-FIRST INPUT SELECTION\n_new_cache_was_built = False\n# ============================================================================\n# Experiment 1 uses the ORIGINAL 0.27 baseline architecture, so we want the\n# compatible wide-v2 cache from the previous run:\n#   cached_specs_wide_v2.npy\n#   cached_eeg_wide_v2.npy\n#   cached_labels_v2.npy\n#\n# These are the exact cache files shown in the previous Kaggle run. If the\n# cache dataset is attached under /kaggle/input, load it directly and DO NOT\n# rebuild the expensive caches. The v4 CWT cache is a different representation\n# and is intentionally ignored by this baseline experiment.\n\nfrom pathlib import Path\nimport hashlib\nimport re\n\nCACHE_SEARCH_ROOTS = ['/kaggle/input', '/kaggle/working']\n\n# PAPER-SAFE CACHE POLICY:\n# The paper experiment uses ONE canonical cache location in /kaggle/working.\n# We do not reuse /kaggle/input cache files directly because an attached legacy\n# cache has no provenance manifest. On the first paper-safe run, the expensive\n# cache is rebuilt into /kaggle/working and fingerprinted. Subsequent modality\n# runs reuse that fingerprinted working cache.\nCACHE_X_PATH = '/kaggle/working/cached_specs_wide_v2.npy'\nCACHE_EEG_PATH = '/kaggle/working/cached_eeg_wide_v2.npy'\nCACHE_Y_PATH = '/kaggle/working/cached_labels_v2.npy'\nCACHE_FINGERPRINT_PATH = os.path.join(\n    '/kaggle/working',\n    'hms_cache_sample_fingerprint.sha256',\n)\n_new_cache_was_built = False\n\nFINGERPRINT_COLUMNS = [\n    'patient_id',\n    'eeg_id',\n    'spectrogram_id',\n    'eeg_label_offset_seconds',\n    'spectrogram_label_offset_seconds',\n]\n\nmissing_fp_cols = [c for c in FINGERPRINT_COLUMNS if c not in agg_df.columns]\nif missing_fp_cols:\n    raise RuntimeError(\n        'Cannot verify cache identity; missing columns: ' + repr(missing_fp_cols)\n    )\n\ndef _cache_valid_at(x_path, eeg_path, y_path):\n    try:\n        y = np.load(y_path, mmap_mode='r')\n        x = np.load(x_path, mmap_mode='r')\n        e = np.load(eeg_path, mmap_mode='r')\n        return (\n            len(y) == len(agg_df) and\n            x.shape == (len(agg_df), SPEC_TARGET_H, WIDE_W) and\n            e.shape == (len(agg_df), 16, EEG_WIDE_LEN) and\n            y.shape[1] == len(vote_cols)\n        )\n    except Exception as exc:\n        print(f'Cache validation failed for {x_path}: {exc}')\n        return False\n\ndef _canonicalize_identity_value(value):\n    if pd.isna(value):\n        return '<NA>'\n    if isinstance(value, (np.integer, int)):\n        return str(int(value))\n    if isinstance(value, (np.floating, float)):\n        return format(float(value), '.12g')\n    return str(value)\n\ndef compute_sample_fingerprint(frame):\n    h = hashlib.sha256()\n    for row in frame[FINGERPRINT_COLUMNS].itertuples(index=False, name=None):\n        payload = '\\x1f'.join(\n            _canonicalize_identity_value(v) for v in row\n        ).encode('utf-8')\n        h.update(payload)\n        h.update(b'\\n')\n    return h.hexdigest()\n\nCURRENT_SAMPLE_FINGERPRINT = compute_sample_fingerprint(agg_df)\n\n# IMPORTANT: never select an arbitrary cache from /kaggle/input for the paper\n# experiment. An attached legacy cache may be shape-compatible but unverified.\n# The only cache that may be reused is the fingerprinted working cache created\n# by this notebook.\nexisting_working_cache = _cache_valid_at(\n    CACHE_X_PATH, CACHE_EEG_PATH, CACHE_Y_PATH\n)\nfingerprint_exists = os.path.exists(CACHE_FINGERPRINT_PATH)\n\nif fingerprint_exists:\n    if not existing_working_cache:\n        raise RuntimeError(\n            '\\nFINGERPRINTED CACHE FILES ARE MISSING OR INVALID\\n'\n            'The fingerprint manifest exists, but the corresponding working '\n            'cache files are missing or have incompatible shapes. Rebuild the '\n            'cache by removing the stale fingerprint and rerunning this cell.'\n        )\n\n    with open(CACHE_FINGERPRINT_PATH, 'r', encoding='utf-8') as f:\n        stored_text = f.read().strip()\n\n    match = re.search(\n        r'^sha256=([0-9a-f]{64})$',\n        stored_text,\n        flags=re.MULTILINE,\n    )\n    if match is None:\n        raise RuntimeError(\n            f'Invalid cache fingerprint file: {CACHE_FINGERPRINT_PATH}'\n        )\n\n    stored_fingerprint = match.group(1)\n    if stored_fingerprint != CURRENT_SAMPLE_FINGERPRINT:\n        raise RuntimeError(\n            '\\nCACHE IDENTITY MISMATCH\\n'\n            'The pre-existing working cache does not correspond to the current '\n            'labeled-window identity/order.\\n\\n'\n            f'Stored fingerprint:  {stored_fingerprint}\\n'\n            f'Current fingerprint: {CURRENT_SAMPLE_FINGERPRINT}\\n\\n'\n            'Rebuild the cache before training.'\n        )\n\n    print('CACHE IDENTITY CHECK PASSED')\n    print(f'  fingerprint: {CURRENT_SAMPLE_FINGERPRINT}')\n    print(f'  rows: {len(agg_df)}')\n\n    cached_specs = np.load(CACHE_X_PATH, mmap_mode='r')\n    cached_eeg = np.load(CACHE_EEG_PATH, mmap_mode='r')\n    cached_labels = np.load(CACHE_Y_PATH, mmap_mode='r')\nelse:\n    # No manifest means the working cache, if present, is legacy/unverified.\n    # Rebuild it in-place rather than touching or trusting /kaggle/input caches.\n    print('\\nNO VERIFIED PAPER CACHE FOUND')\n    print('Any pre-existing unverified v2 cache will be ignored and rebuilt in /kaggle/working.')\n\n    n = len(agg_df)\n    print(\n        f'\\nBuilding paper-safe caches for {n:,} labeled windows: '\n        f'spectrogram {SPEC_TARGET_H}x{WIDE_W}, EEG {EEG_WIDE_LEN} samples...'\n    )\n\n    cached_specs = np.lib.format.open_memmap(\n        CACHE_X_PATH,\n        mode='w+',\n        dtype=np.float16,\n        shape=(n, SPEC_TARGET_H, WIDE_W),\n    )\n    cached_eeg = np.lib.format.open_memmap(\n        CACHE_EEG_PATH,\n        mode='w+',\n        dtype=np.float16,\n        shape=(n, 16, EEG_WIDE_LEN),\n    )\n    cached_labels = agg_df[vote_cols].values.astype(np.float32)\n    _new_cache_was_built = True\n\n    spec_cache_by_id = {}\n    eeg_cache_by_id = {}\n\n    for i, row in agg_df.iterrows():\n        sid = row['spectrogram_id']\n        if sid not in spec_cache_by_id:\n            spec_cache_by_id.clear()\n            spec_cache_by_id[sid] = pd.read_parquet(\n                f'{spec_path}/{sid}.parquet'\n            )\n        spec_raw = spec_cache_by_id[sid]\n        cached_specs[i] = build_spectrogram_image(\n            spec_raw,\n            row['spectrogram_label_offset_seconds'],\n        )\n\n        eid = row['eeg_id']\n        if eid not in eeg_cache_by_id:\n            eeg_cache_by_id.clear()\n            eeg_cache_by_id[eid] = pd.read_parquet(\n                f'{eeg_path}/{eid}.parquet'\n            )\n        eeg_raw = eeg_cache_by_id[eid]\n        cached_eeg[i] = build_eeg_window(\n            eeg_raw,\n            row['eeg_label_offset_seconds'],\n        )\n\n    cached_specs.flush()\n    cached_eeg.flush()\n    np.save(CACHE_Y_PATH, cached_labels)\n    cached_labels = np.load(CACHE_Y_PATH, mmap_mode='r')\n\n    with open(CACHE_FINGERPRINT_PATH, 'w', encoding='utf-8') as f:\n        f.write(\n            'sha256=' + CURRENT_SAMPLE_FINGERPRINT + '\\n'\n            f'rows={len(agg_df)}\\n'\n            'columns=' + ','.join(FINGERPRINT_COLUMNS) + '\\n'\n        )\n\n    print('CACHE BUILT BY THIS NOTEBOOK')\n    print('Wrote cache fingerprint.')\n    print(f'  fingerprint: {CURRENT_SAMPLE_FINGERPRINT}')\n    print(f'  rows: {len(agg_df)}')\n\n# IMPORTANT: cached_labels_v2.npy may have been created by an earlier smoothed-target\n# run. For this no-smoothing ablation, specs/EEG can be reused, but the LABEL ARRAY\n# must come from the current raw vote proportions built above. This keeps the cache\n# for expensive inputs while preventing a silent target-definition mismatch.\n_expected_raw_labels = agg_df[vote_cols].to_numpy(dtype=np.float32)\nif len(cached_specs) != len(_expected_raw_labels) or len(cached_eeg) != len(_expected_raw_labels):\n    raise RuntimeError(\n        f'Cache/sample count mismatch: specs={len(cached_specs)}, '\n        f'eeg={len(cached_eeg)}, raw_labels={len(_expected_raw_labels)}. '\n        'Do not reuse this cache for the current target ordering.'\n    )\ncached_labels = _expected_raw_labels\nassert np.allclose(cached_labels, raw_vote_targets, atol=1e-7), \\\n    'Internal target mismatch: cached_labels and raw_vote_targets differ.'\nprint('No-smoothing target override: using current raw vote proportions for cached_labels.')\nprint('  cached_labels shape:', cached_labels.shape)\nprint('  max |cached_labels - raw_vote_targets|:', float(np.max(np.abs(cached_labels - raw_vote_targets))))\n\n# Expose the canonical names used by the rest of the notebook.\nCACHE_X = CACHE_X_PATH\nCACHE_EEG = CACHE_EEG_PATH\nCACHE_Y = CACHE_Y_PATH\n","metadata":{"execution":{"iopub.status.busy":"2026-09-11T06:03:14.862204Z","iopub.execute_input":"2026-09-11T06:03:14.862689Z","iopub.status.idle":"2026-09-11T06:17:59.865001Z","shell.execute_reply.started":"2026-09-11T06:03:14.862665Z","shell.execute_reply":"2026-09-11T06:17:59.864271Z"},"papermill":{"duration":55.672251,"end_time":"2026-08-22T14:49:36.4662+00:00","exception":false,"start_time":"2026-08-22T14:48:40.793949+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Cached samples: {len(cached_specs)}\")\nprint(f\"Spectrogram shape: {cached_specs[0].shape}\")\nprint(f\"Label shape: {cached_labels[0].shape}\")\nprint(f\"Label sum (should be ~1.0): {cached_labels[0].sum():.4f}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-09-11T06:17:59.866341Z","iopub.execute_input":"2026-09-11T06:17:59.866679Z","iopub.status.idle":"2026-09-11T06:17:59.872119Z","shell.execute_reply.started":"2026-09-11T06:17:59.866655Z","shell.execute_reply":"2026-09-11T06:17:59.871382Z"},"id":"239LOisQ_8sI","outputId":"438bccc4-3b28-4e16-d27d-74b6f2820293","papermill":{"duration":0.011912,"end_time":"2026-08-22T14:49:36.483365+00:00","exception":false,"start_time":"2026-08-22T14:49:36.471453+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader, WeightedRandomSampler\nfrom sklearn.model_selection import GroupKFold\n\nNUM_EPOCHS = 15\nN_FOLDS = 10\n\n# Number of folds to ACTUALLY train. Set to 1 for a quick experiment.\nFOLDS_TO_RUN = 10\n\nif not (1 <= FOLDS_TO_RUN <= N_FOLDS):\n    raise ValueError(f\"FOLDS_TO_RUN must be between 1 and {N_FOLDS}, got {FOLDS_TO_RUN}\")\n\ngkf = GroupKFold(n_splits=N_FOLDS)\nsplits = list(gkf.split(agg_df, groups=agg_df['patient_id']))\n\n# IMPORTANT: this is the single source of truth used by the training loop.\n# It prevents later code from accidentally iterating over all configured splitss.\nrun_splits = splits[:FOLDS_TO_RUN]\nRUN_FOLD_IDS = list(range(len(run_splits)))\n\nprint(f\"Generated {len(splits)} total folds; TRAINING ONLY {len(run_splits)} fold(s): {RUN_FOLD_IDS}\")\nfor fold, (train_idx, val_idx) in zip(RUN_FOLD_IDS, run_splits):\n    train_patients = set(agg_df.iloc[train_idx]['patient_id'])\n    val_patients = set(agg_df.iloc[val_idx]['patient_id'])\n    overlap = len(train_patients & val_patients)\n    print(f\"  Fold {fold + 1}/{N_FOLDS}: train={len(train_idx)} val={len(val_idx)} \"\n          f\"patient_overlap={overlap} (should be 0)\")\n","metadata":{"id":"qrf1YAUhCjtR","outputId":"595b10a7-d0c4-4b66-aad3-af44bb9e09af","papermill":{"duration":1.374296,"end_time":"2026-08-22T14:49:37.862868+00:00","exception":false,"start_time":"2026-08-22T14:49:36.488572+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# MODEL CONFIGURATION\n# ============================================================================\n# Change ONLY this variable between the three main experiments:\n#   \"spectrogram\" -> spectrogram-only model\n#   \"eeg\"         -> raw-EEG-only model\n#   \"multimodal\"  -> spectrogram + raw-EEG fusion model\n#\n# Everything else is held constant across these runs.\nMODALITY = \"spectrogram\"\n\n# Keep the rest of the current experiment settings unchanged.\nDROPOUT_RATE = 0.50\n\nimport torch\nimport timm\nimport torch.nn as nn\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n\nclass EEG1DBranch(nn.Module):\n    \"\"\"Compact 1D CNN over the 16-channel bipolar EEG montage.\"\"\"\n\n    def __init__(self, in_channels=16, out_dim=128):\n        super().__init__()\n\n        def block(c_in, c_out, k=7, stride=2):\n            return nn.Sequential(\n                nn.Conv1d(c_in, c_out, kernel_size=k, stride=stride, padding=k // 2),\n                nn.BatchNorm1d(c_out),\n                nn.ReLU(inplace=True),\n            )\n\n        self.net = nn.Sequential(\n            block(in_channels, 32),\n            block(32, 64),\n            block(64, 128),\n            block(128, 128),\n            block(128, out_dim),\n        )\n        self.pool = nn.AdaptiveAvgPool1d(1)\n        self.out_dim = out_dim\n\n    def forward(self, x):\n        x = self.net(x)\n        x = self.pool(x).squeeze(-1)\n        return x\n\n\nclass HMSSingleOrFusionModel(nn.Module):\n    \"\"\"HMS classifier with one modality switch.\n\n    spectrogram:\n        EfficientNet-B0 over the spectrogram.\n    eeg:\n        Small 1D CNN over the middle-10s raw EEG.\n    multimodal:\n        Both branches, followed by feature concatenation.\n\n    The forward signature stays (spec, eeg) for compatibility with the\n    existing dataset/training/evaluation code.\n    \"\"\"\n\n    def __init__(\n        self,\n        modality=MODALITY,\n        spec_model_name=\"tf_efficientnet_b0_ns\",\n        pretrained=True,\n        drop_rate=DROPOUT_RATE,\n        eeg_out_dim=128,\n        num_classes=6,\n    ):\n        super().__init__()\n\n        if modality not in {\"spectrogram\", \"eeg\", \"multimodal\"}:\n            raise ValueError(\n                f\"MODALITY must be 'spectrogram', 'eeg', or 'multimodal'; got {modality!r}\"\n            )\n\n        self.modality = modality\n        self.spec_branch = None\n        self.eeg_branch = None\n\n        if modality in {\"spectrogram\", \"multimodal\"}:\n            self.spec_branch = timm.create_model(\n                spec_model_name,\n                pretrained=pretrained,\n                in_chans=3,\n                num_classes=0,\n                drop_rate=drop_rate,\n            )\n            spec_feat_dim = self.spec_branch.num_features\n\n        if modality in {\"eeg\", \"multimodal\"}:\n            self.eeg_branch = EEG1DBranch(\n                in_channels=16,\n                out_dim=eeg_out_dim,\n            )\n\n        if modality == \"spectrogram\":\n            classifier_in_dim = spec_feat_dim\n        elif modality == \"eeg\":\n            classifier_in_dim = eeg_out_dim\n        else:\n            classifier_in_dim = spec_feat_dim + eeg_out_dim\n\n        self.classifier = nn.Sequential(\n            nn.Linear(classifier_in_dim, 256),\n            nn.ReLU(inplace=True),\n            nn.Dropout(drop_rate),\n            nn.Linear(256, num_classes),\n        )\n\n    def forward(self, spec, eeg):\n        features = []\n\n        if self.modality in {\"spectrogram\", \"multimodal\"}:\n            features.append(self.spec_branch(spec))\n\n        if self.modality in {\"eeg\", \"multimodal\"}:\n            features.append(self.eeg_branch(eeg))\n\n        x = features[0] if len(features) == 1 else torch.cat(features, dim=1)\n        return self.classifier(x)\n\n\ndef build_model(\n    modality=MODALITY,\n    spec_model_name=\"tf_efficientnet_b0_ns\",\n    pretrained=True,\n    drop_rate=DROPOUT_RATE,\n):\n    model = HMSSingleOrFusionModel(\n        modality=modality,\n        spec_model_name=spec_model_name,\n        pretrained=pretrained,\n        drop_rate=drop_rate,\n    )\n    model.to(device)\n    return model\n\n\ndef set_backbone_trainable(model, trainable):\n    # Keep the classifier head trainable; freeze/unfreeze whichever modality\n    # backbones are actually present in this run.\n    for name, param in model.named_parameters():\n        if \"classifier\" not in name:\n            param.requires_grad = trainable\n\n\nprint(f\"Device: {device}\")\nprint(f\"MODALITY = {MODALITY}\")\nprint(f\"DROPOUT_RATE = {DROPOUT_RATE:.2f}\")\n\n_sanity_model = build_model()\nprint(f\"Model initialized successfully: {MODALITY}\")\ndel _sanity_model\n","metadata":{"id":"XKqh191jDGdA","outputId":"cc059ea8-cdf4-49da-d72b-87b9977c70f2","papermill":{"duration":12.35473,"end_time":"2026-08-22T14:49:50.223434+00:00","exception":false,"start_time":"2026-08-22T14:49:37.868704+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Leakage / cache integrity safeguards\n\nThe main modality experiment keeps the patient-grouped splits, targets, augmentation, optimizer, schedule, and sampling policy fixed. Precomputed EEG/spectrogram inputs are reused only when a SHA-256 fingerprint proves their exact sample identity/order matches the current labeled-window table. Missing or mismatched fingerprints stop the run.\n","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\n\n# Primary objective: compare predicted probability distributions with the\n# expert vote probability distributions.\ncriterion = nn.KLDivLoss(reduction=\"batchmean\")\n\n# IMPORTANT FOR THE RESEARCH QUESTION:\n# Keep this OFF for the main experiment. The expert targets are soft\n# probability distributions, so adding a hard-label auxiliary loss would push\n# the model toward the argmax class instead of preserving uncertainty.\nUSE_AUX_ENTMAX = False\nAUX_ENTMAX_WEIGHT = 0.0\n\ndef build_optimiser_scheduler(model, epochs=NUM_EPOCHS):\n    optimiser = optim.AdamW(\n        model.parameters(),\n        lr=5e-4,\n        weight_decay=1e-2,\n    )\n\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(\n        optimiser,\n        T_max=epochs,\n        eta_min=1e-6,\n    )\n    return optimiser, scheduler\n\nprint(\"Loss: KL divergence with standard softmax outputs\")\nprint(f\"Auxiliary entmax enabled: {USE_AUX_ENTMAX}\")\n","metadata":{"id":"C8wwBUiBEueh","papermill":{"duration":0.017927,"end_time":"2026-08-22T14:49:50.247133+00:00","exception":false,"start_time":"2026-08-22T14:49:50.229206+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.amp import autocast, GradScaler\n\nMIXUP_ALPHA = 0.2\nMIXUP_PROB = 0.5\n\nUSE_AMP = (device.type == \"cuda\")\nscaler = GradScaler(device.type if USE_AMP else \"cpu\", enabled=USE_AMP)\n\n\ndef mixup_batch(specs, eeg, labels, alpha=MIXUP_ALPHA):\n    lam = np.random.beta(alpha, alpha)\n    perm = torch.randperm(specs.size(0), device=specs.device)\n\n    mixed_specs = lam * specs + (1 - lam) * specs[perm]\n    mixed_eeg = lam * eeg + (1 - lam) * eeg[perm]\n    mixed_labels = lam * labels + (1 - lam) * labels[perm]\n\n    return mixed_specs, mixed_eeg, mixed_labels\n\n\ndef train_one_epoch(model, loader, optimiser, criterion, device):\n    model.train()\n    total_kl = 0.0\n    n_batches = len(loader)\n\n    for batch_idx, (specs, eeg, labels) in enumerate(loader):\n        specs = specs.to(device, non_blocking=True)\n        eeg = eeg.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        if np.random.rand() < MIXUP_PROB:\n            specs, eeg, labels = mixup_batch(specs, eeg, labels)\n\n        optimiser.zero_grad(set_to_none=True)\n\n        with autocast(device.type, enabled=USE_AMP):\n            logits = model(specs, eeg)\n\n            # Standard softmax probability distribution + KL divergence.\n            # This directly matches the six-class expert vote distribution.\n            log_probability = F.log_softmax(logits.float(), dim=1)\n            kl_loss = criterion(log_probability, labels)\n\n            # Main experiment objective = KL ONLY.\n            backward_loss = kl_loss\n\n        scaler.scale(backward_loss).backward()\n        scaler.step(optimiser)\n        scaler.update()\n\n        total_kl += kl_loss.item()\n\n        if batch_idx % 100 == 0:\n            print(\n                f\"Batch {batch_idx}/{n_batches} - KL Loss: {kl_loss.item():.4f}\"\n            )\n\n    return total_kl / max(n_batches, 1)\n\n\ndef validate(model, loader, criterion, device):\n    \"\"\"Return the exact per-sample validation KL against expert vote distributions.\"\"\"\n    model.eval()\n    total_kl = 0.0\n    total_n = 0\n\n    with torch.no_grad():\n        for specs, eeg, labels in loader:\n            specs = specs.to(device, non_blocking=True)\n            eeg = eeg.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            logits = model(specs, eeg)\n            log_probs = F.log_softmax(logits.float(), dim=1)\n            loss = criterion(log_probs, labels)\n\n            bs = labels.size(0)\n            total_kl += float(loss.item()) * bs\n            total_n += bs\n\n    return total_kl / max(total_n, 1)\n","metadata":{"papermill":{"duration":0.019889,"end_time":"2026-08-22T14:49:50.272709+00:00","exception":false,"start_time":"2026-08-22T14:49:50.25282+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# EXPERIMENT CONFIGURATION\n# ============================================================================\n# The modality is the primary experimental variable.\n#\n# Weighted sampling is kept ON across all three modality runs so class\n# imbalance is controlled identically and does not become a modality-specific\n# factor. It is NOT the research variable here.\nUSE_WEIGHTED_SAMPLING = True\n\nprint(\"=\" * 70)\nprint(\"MAIN RESEARCH EXPERIMENT\")\nprint(\"=\" * 70)\nprint(f\"MODALITY: {MODALITY}\")\nprint(f\"Weighted sampling: {USE_WEIGHTED_SAMPLING}\")\nprint(f\"Auxiliary entmax: {USE_AUX_ENTMAX}\")\nprint(f\"Dropout: {DROPOUT_RATE:.2f}\")\nprint(f\"Mixup alpha/probability: {MIXUP_ALPHA} / {MIXUP_PROB}\")\nprint(\"Primary objective: soft-target KL divergence\")\nprint(\"=\" * 70)\n","metadata":{"papermill":{"duration":0.012918,"end_time":"2026-08-22T14:49:50.291323+00:00","exception":false,"start_time":"2026-08-22T14:49:50.278405+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report, confusion_matrix\n\nCLASS_NAMES = [c.replace('_vote', '') for c in vote_cols]\n\n\ndef compute_accuracy(model, loader, device):\n    model.eval()\n    correct, total = 0, 0\n    with torch.no_grad():\n        for specs, eeg, labels in loader:\n            specs, eeg, labels = specs.to(device), eeg.to(device), labels.to(device)\n            logits = model(specs, eeg)\n            preds = logits.argmax(dim=1)\n            true = labels.argmax(dim=1)\n            correct += (preds == true).sum().item()\n            total += labels.size(0)\n    return correct / total\n\n\ndef print_class_report(true, preds, class_names=CLASS_NAMES):\n    \"\"\"Per-class precision/recall/F1 + confusion matrix. Overall accuracy\n    alone hides which classes the model is actually failing on when classes\n    are imbalanced (see cell 4's value_counts).\"\"\"\n    print(classification_report(true, preds, target_names=class_names, digits=3, zero_division=0))\n    cm = confusion_matrix(true, preds)\n    print(\"Confusion matrix (rows=true, cols=pred):\")\n    print(\"           \" + \" \".join(f\"{c[:6]:>6}\" for c in class_names))\n    for name, row in zip(class_names, cm):\n        print(f\"{name[:10]:>10} \" + \" \".join(f\"{v:>6d}\" for v in row))\n","metadata":{"papermill":{"duration":0.014404,"end_time":"2026-08-22T14:49:50.311567+00:00","exception":false,"start_time":"2026-08-22T14:49:50.297163+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.005474,"end_time":"2026-08-22T14:49:50.3227+00:00","exception":false,"start_time":"2026-08-22T14:49:50.317226+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main modality comparison\n\nChange only `MODALITY`:\n\n- `\"spectrogram\"` = spectrogram-only\n- `\"eeg\"` = raw-EEG-only\n- `\"multimodal\"` = spectrogram + raw EEG\n\nFor the primary comparison, all three arms receive the same 15-epoch budget\nand all available backbones are trainable from epoch 0. This avoids giving\nthe pretrained spectrogram branch a special warmup while freezing the\nrandomly initialized EEG branch.\n\nAll other settings remain fixed: raw expert-vote targets, KL objective,\nweighted sampler, dropout, mixup, optimizer, schedule, and patient-level\nGroupKFold splits.\n\nThe auxiliary entmax hard-label loss is disabled. The model outputs a\nstandard softmax probability distribution, and KL divergence to the expert\nvote distribution is the primary research metric.\n","metadata":{"papermill":{"duration":0.005264,"end_time":"2026-08-22T14:49:50.333524+00:00","exception":false,"start_time":"2026-08-22T14:49:50.32826+00:00","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# FINAL-COMPARISON PROTOCOL\n\nFor the final paper comparison, rerun **all three conditions** under this same\ntraining protocol:\n\n- `MODALITY = \"spectrogram\"`\n- `MODALITY = \"eeg\"`\n- `MODALITY = \"multimodal\"`\n\nAll available backbones are trainable from epoch 0. Do not mix these results\nwith the earlier freeze-based runs.\n\nKeep the same raw expert-vote targets, KL objective, weighted sampler,\ndropout, mixup, optimizer, schedule, and patient-level GroupKFold splits.\n\nThe primary research metric is KL divergence between the expert vote\ndistribution and the model's softmax probability distribution.\n","metadata":{}},{"cell_type":"code","source":"import copy\n\nfold_results = []\nfold_histories = []\noof_preds, oof_true, oof_probs = [], [], []\n\nBATCH_SIZE = 32\nCHECKPOINT_PREFIX = f\"best_model_{MODALITY}_fold\"\n\nfor fold, (train_idx, val_idx) in zip(RUN_FOLD_IDS, run_splits):\n    print(f\"\\n{'='*60}\\nFOLD {fold + 1}/{N_FOLDS} [{MODALITY}]\\n{'='*60}\")\n\n    train_dataset = HMSFusionDataset(\n        cached_specs, cached_eeg, cached_labels, train_idx, train=True\n    )\n    val_dataset = HMSFusionDataset(\n        cached_specs, cached_eeg, cached_labels, val_idx, train=False\n    )\n\n    sampler = None\n    class_counts = None\n    class_sampling_weights = None\n\n    if USE_WEIGHTED_SAMPLING:\n        # Used only for sampling; the target remains the full soft vote vector.\n        train_hard_classes = cached_labels[train_idx].argmax(axis=1)\n\n        class_counts = np.bincount(\n            train_hard_classes,\n            minlength=len(CLASS_NAMES),\n        ).astype(np.float64)\n\n        class_sampling_weights = 1.0 / np.sqrt(\n            np.maximum(class_counts, 1.0)\n        )\n        sample_weights = class_sampling_weights[train_hard_classes]\n\n        sampler = WeightedRandomSampler(\n            weights=torch.as_tensor(sample_weights, dtype=torch.double),\n            num_samples=len(train_dataset),\n            replacement=True,\n        )\n\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=BATCH_SIZE,\n        sampler=sampler,\n        shuffle=(sampler is None),\n        num_workers=2,\n        persistent_workers=True,\n    )\n\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=2,\n        persistent_workers=True,\n    )\n\n    print(f\"Train samples: {len(train_dataset)} | Val samples: {len(val_dataset)}\")\n\n    if USE_WEIGHTED_SAMPLING:\n        print(\"Training class counts / sampler weights:\")\n        for i, name in enumerate(CLASS_NAMES):\n            print(\n                f\"  {name:8s}: n={int(class_counts[i]):6d} \"\n                f\"weight={class_sampling_weights[i]:.6f}\"\n            )\n    else:\n        print(\"Weighted sampling disabled; training uses natural sample frequency.\")\n\n    model = build_model(modality=MODALITY)\n    optimiser, scheduler = build_optimiser_scheduler(model)\n\n    best_val_loss = float(\"inf\")\n    best_model_state = None\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_acc\": []}\n\n    # FAIR MODALITY COMPARISON:\n    # All available backbones are trainable from epoch 0.\n    # This removes the artificial disadvantage of freezing the randomly\n    # initialized EEG branch during the first epoch.\n    set_backbone_trainable(model, True)\n    print(\"  [all available backbones trainable from epoch 0]\")\n\n    for epoch in range(NUM_EPOCHS):\n        print(\n            f\"\\nEpoch {epoch + 1}/{NUM_EPOCHS} | \"\n            f\"LR: {scheduler.get_last_lr()[0]:.2e}\"\n        )\n\n        train_loss = train_one_epoch(\n            model, train_loader, optimiser, criterion, device\n        )\n        val_loss = validate(model, val_loader, criterion, device)\n        val_acc = compute_accuracy(model, val_loader, device)\n        scheduler.step()\n\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_acc\"].append(val_acc)\n\n        print(\n            f\"Train Loss: {train_loss:.4f} | \"\n            f\"Val KL: {val_loss:.4f} | \"\n            f\"Val Acc: {val_acc:.4f}\"\n        )\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            best_model_state = copy.deepcopy(model.state_dict())\n\n            ckpt_path = f\"{CHECKPOINT_PREFIX}{fold}.pt\"\n            torch.save(best_model_state, ckpt_path)\n\n            print(\n                f\"  -> new best model saved for fold {fold + 1}: \"\n                f\"{ckpt_path} (val_KL={val_loss:.4f})\"\n            )\n\n    model.load_state_dict(best_model_state)\n\n    best_epoch_idx = int(np.argmin(history[\"val_loss\"]))\n    best_val_acc = history[\"val_acc\"][best_epoch_idx]\n\n    print(\n        f\"\\nFold {fold + 1} done. \"\n        f\"Best val_KL={best_val_loss:.4f} \"\n        f\"(epoch {best_epoch_idx + 1}), \"\n        f\"val_acc={best_val_acc:.4f}\"\n    )\n\n    model.eval()\n    fold_probs = []\n\n    with torch.no_grad():\n        for specs, eeg, labels in val_loader:\n            specs = specs.to(device)\n            eeg = eeg.to(device)\n            labels = labels.to(device)\n\n            logits = model(specs, eeg)\n            probs = F.softmax(logits.float(), dim=1).cpu().numpy()\n\n            fold_probs.append(probs)\n            oof_preds.extend(probs.argmax(axis=1).tolist())\n            oof_true.extend(labels.argmax(dim=1).cpu().tolist())\n\n    fold_probs = np.concatenate(fold_probs, axis=0)\n    fold_raw_targets = raw_vote_targets[val_idx]\n\n    target_gap = float(\n        np.max(\n            np.abs(\n                np.asarray(cached_labels[val_idx], dtype=np.float32)\n                - np.asarray(fold_raw_targets, dtype=np.float32)\n            )\n        )\n    )\n\n    print(\n        f\"  Target alignment max|cached_labels - raw_votes| = \"\n        f\"{target_gap:.10f}\"\n    )\n\n    if target_gap > 1e-7:\n        raise RuntimeError(\n            f\"Cached-label/raw-vote target mismatch in fold {fold+1}: \"\n            f\"max diff={target_gap:.10f}\"\n        )\n\n    # Primary research metric:\n    # KL(expert vote distribution || model probability distribution).\n    fold_raw_kl = np.sum(\n        np.where(\n            fold_raw_targets > 0,\n            fold_raw_targets\n            * (\n                np.log(np.maximum(fold_raw_targets, 1e-12))\n                - np.log(np.maximum(fold_probs, 1e-12))\n            ),\n            0.0,\n        ),\n        axis=1,\n    ).mean()\n\n    oof_probs.extend(fold_probs.tolist())\n\n    metric_gap = abs(float(best_val_loss) - float(fold_raw_kl))\n    print(\n        f\"  Metric consistency check: \"\n        f\"|val_KL - raw_vote_KL| = {metric_gap:.8f}\"\n    )\n\n    if metric_gap > 5e-6:\n        raise RuntimeError(\n            f\"Metric mismatch detected in fold {fold+1}: \"\n            f\"validation KL={best_val_loss:.6f}, \"\n            f\"raw-vote KL={fold_raw_kl:.6f}\"\n        )\n\n    fold_results.append(\n        {\n            \"fold\": fold,\n            \"modality\": MODALITY,\n            \"best_val_loss\": best_val_loss,\n            \"best_val_acc\": best_val_acc,\n            \"raw_vote_kl\": fold_raw_kl,\n        }\n    )\n    fold_histories.append(history)\n\n    print(f\"  Raw-vote KL: {fold_raw_kl:.4f}\")\n\n    del model, optimiser, scheduler\n    torch.cuda.empty_cache()\n\n\nval_losses = [r[\"best_val_loss\"] for r in fold_results]\nraw_vote_kls = [r[\"raw_vote_kl\"] for r in fold_results]\nval_accs = [r[\"best_val_acc\"] for r in fold_results]\n\nprint(\n    f\"\\n{'='*60}\\n\"\n    f\"CV SUMMARY | {MODALITY} | \"\n    f\"{len(fold_results)}/{N_FOLDS} folds run\\n\"\n    f\"{'='*60}\"\n)\n\nfor r in fold_results:\n    print(\n        f\"  Fold {r['fold'] + 1}: \"\n        f\"val_KL={r['best_val_loss']:.4f} | \"\n        f\"raw_vote_KL={r['raw_vote_kl']:.4f} | \"\n        f\"val_acc={r['best_val_acc']:.4f}\"\n    )\n\nprint(\n    f\"\\nMean CV val_KL: \"\n    f\"{np.mean(val_losses):.4f} +/- {np.std(val_losses):.4f}\"\n)\nprint(\n    f\"Mean CV raw-vote KL: \"\n    f\"{np.mean(raw_vote_kls):.4f} +/- {np.std(raw_vote_kls):.4f}\"\n)\nprint(\n    f\"Mean CV val_acc: \"\n    f\"{np.mean(val_accs):.4f} +/- {np.std(val_accs):.4f}\"\n)\n","metadata":{"papermill":{"duration":33210.449997,"end_time":"2026-08-23T00:03:20.78916+00:00","exception":false,"start_time":"2026-08-22T14:49:50.339163+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import Counter\n\n# NOTE: these are aggregated out-of-fold predictions across all configured folds\n# (each sample's prediction comes from the fold where it was held out) - the\n# correct basis for the paper's reported class distribution / confusion analysis.\nprint(\"OOF predicted class distribution:\", Counter(oof_preds))\nprint(\"OOF true class distribution:\", Counter(oof_true))\nprint()\nprint_class_report(oof_true, oof_preds)\n","metadata":{"papermill":{"duration":0.23241,"end_time":"2026-08-23T00:03:21.100718+00:00","exception":false,"start_time":"2026-08-23T00:03:20.868308+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def ensemble_predict(loader, n_folds=None, device=device):\n    \"\"\"Average probabilities across trained checkpoints for the selected modality.\"\"\"\n\n    if n_folds is None:\n        n_folds = len(run_splits)\n\n    if not (1 <= n_folds <= N_FOLDS):\n        raise ValueError(\n            f\"n_folds must be between 1 and {N_FOLDS}, got {n_folds}\"\n        )\n\n    models_list = []\n\n    for fold in range(n_folds):\n        ckpt = f\"best_model_{MODALITY}_fold{fold}.pt\"\n\n        if not os.path.exists(ckpt):\n            raise FileNotFoundError(\n                f\"Missing {ckpt}. Run the selected modality first or set \"\n                f\"FOLDS_TO_RUN to the number of trained folds.\"\n            )\n\n        m = build_model(modality=MODALITY)\n        m.load_state_dict(torch.load(ckpt, map_location=device))\n        m.eval()\n        models_list.append(m)\n\n    all_probs = []\n\n    with torch.no_grad():\n        for specs, eeg, _ in loader:\n            specs = specs.to(device)\n            eeg = eeg.to(device)\n\n            batch_probs = torch.zeros(specs.size(0), 6, device=device)\n\n            for m in models_list:\n                logits = m(specs, eeg)\n                batch_probs += F.softmax(logits.float(), dim=1)\n\n            batch_probs /= n_folds\n            all_probs.append(batch_probs.cpu())\n\n    for m in models_list:\n        del m\n\n    torch.cuda.empty_cache()\n    return torch.cat(all_probs, dim=0)\n","metadata":{"papermill":{"duration":0.08668,"end_time":"2026-08-23T00:03:21.265752+00:00","exception":false,"start_time":"2026-08-23T00:03:21.179072+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Running the three research conditions\n\nRun the notebook three times, changing only `MODALITY`:\n\n1. `MODALITY = \"spectrogram\"`\n2. `MODALITY = \"eeg\"`\n3. `MODALITY = \"multimodal\"`\n\nCheckpoints include the modality in their filename, so separate runs do not\noverwrite each other.\n\nPrimary comparison:\n**lower raw-vote KL = the model's predicted probability distribution is closer\nto the expert probability distribution.**\n\nRecommended paper table:\n\n| Model | Raw-vote KL ↓ | Accuracy ↑ | Macro F1 ↑ |\n|---|---:|---:|---:|\n| Spectrogram only | ... | ... | ... |\n| Raw EEG only | ... | ... | ... |\n| Multimodal | ... | ... | ... |\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\ndef plot_confusion_matrix(true, preds, class_names=CLASS_NAMES, save_path='/kaggle/working/confusion_matrix.png'):\n    \"\"\"Side-by-side raw-count and row-normalized (recall) confusion matrices\n    for the aggregated out-of-fold predictions. Row-normalized view matters\n    here specifically because 'other' dominates the class distribution -\n    raw counts alone mostly just show class size, not which classes the\n    model actually confuses.\"\"\"\n    cm_raw = confusion_matrix(true, preds)\n    cm_norm = confusion_matrix(true, preds, normalize='true')\n\n    fig, axes = plt.subplots(1, 2, figsize=(13, 5.5))\n\n    for ax, cm, title, fmt, vmax in [\n        (axes[0], cm_raw, 'Counts', 'd', None),\n        (axes[1], cm_norm, 'Normalized (recall)', '.2f', 1.0),\n    ]:\n        im = ax.imshow(cm, cmap='Blues', vmin=0, vmax=vmax)\n        ax.set_xticks(range(len(class_names)))\n        ax.set_yticks(range(len(class_names)))\n        ax.set_xticklabels(class_names, rotation=45, ha='right')\n        ax.set_yticklabels(class_names)\n        ax.set_xlabel('Predicted')\n        ax.set_ylabel('True')\n        ax.set_title(title)\n\n        thresh = cm.max() / 2 if vmax is None else vmax / 2\n        for i in range(cm.shape[0]):\n            for j in range(cm.shape[1]):\n                val = cm[i, j]\n                ax.text(j, i, format(val, fmt),\n                        ha='center', va='center',\n                        color='white' if val > thresh else 'black',\n                        fontsize=9)\n        fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)\n\n    fig.suptitle(f'OOF Confusion Matrix (aggregated across {len(run_splits)}/{N_FOLDS} folds)')\n    fig.tight_layout()\n    if save_path:\n        fig.savefig(save_path, dpi=200, bbox_inches='tight')\n        print(f\"Saved to {save_path}\")\n    plt.show()\n\n\nplot_confusion_matrix(oof_true, oof_preds)","metadata":{"papermill":{"duration":1.274808,"end_time":"2026-08-23T00:03:22.615948+00:00","exception":false,"start_time":"2026-08-23T00:03:21.34114+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}