{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"cell-117d5c","cell_type":"code","source":"# 2D CNN — EfficientNet-B0 on HMS spectrograms (WORKSHOP VERSION, self-contained)\n# Single-stage training, workshop-scale data (600 train / 120 val rows)\n#\n# Rationale for no 2-step: 2-step training rescues minority-class knowledge from a large,\n# class-imbalanced dataset (Seizure ~256 vs Other ~3050 in the full data). The workshop subset\n# is already balanced (100 rows/class train, 20 rows/class val), so there is no imbalance to\n# correct for — a single training stage is sufficient here.\n#\n# This notebook is fully self-contained — no external .py imports. All classes and functions\n# are defined inline below so you can read, tweak, and re-run any part of the pipeline directly.\n!pip install timm -q","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell-e38331","cell_type":"code","source":"import os, random, time\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import f1_score\nimport timm\n\nIS_KAGGLE = os.path.exists('/kaggle')\nprint(f\"Environment: {'Kaggle' if IS_KAGGLE else 'Local'} | torch {torch.__version__}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell-62ab2c","cell_type":"code","source":"# ============ Config ============\nfrom dataclasses import dataclass, field\n\n@dataclass\nclass Config2D:\n    # ── Data paths ──────────────────────────────────────────────\n    spectrogram_dir: str = \"train_spectrograms\"\n    spec_cache_dir:  str = \"/kaggle/working/spec_cache\"\n\n    # ── Spectrogram image dimensions ────────────────────────────────────────\n    img_height: int = 100   # freq bins per chain (4 chains stacked → 400 total rows)\n    img_width:  int = 300   # time columns per window (2-s resolution → 300 = 600 s)\n\n    # ── Model ───────────────────────────────────────────────────────\n    backbone:    str  = \"efficientnet_b0\"\n    pretrained:  bool = True\n    num_classes: int  = 6\n\n    # ── Training ──────────────────────────────────────────────\n    batch_size:   int   = 32\n    lr:           float = 1e-3\n    weight_decay: float = 1e-4\n    drop_rate:    float = 0.3\n\n    # ── Early stopping ───────────────────────────────────────────\n    patience: int = 10\n\n    # ── Misc ──────────────────────────────────────────────────────\n    seed:        int = 42\n    num_workers: int = 2\n\n    device: str = field(init=False)\n\n    def __post_init__(self):\n        self.device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n\ncfg = Config2D()\n\nif IS_KAGGLE:\n    DATA_ROOT       = '/kaggle/input/competitions/hms-harmful-brain-activity-classification'\n    RAW_TRAIN_PATH  = os.path.join(DATA_ROOT, 'train.csv')\n    SAMPLE_IDS_PATH = '/kaggle/input/datasets/xiaosufrankhu/midas-summer-academy-wk3-eeg/workshop_sample_ids.csv'\n    cfg.spectrogram_dir = os.path.join(DATA_ROOT, 'train_spectrograms')\nelse:\n    DATA_ROOT       = os.path.abspath('../')\n    RAW_TRAIN_PATH  = os.path.abspath('../data_raw/train.csv')\n    SAMPLE_IDS_PATH = os.path.abspath('../data_raw/workshop_sample_ids.csv')\n    cfg.spectrogram_dir = os.path.join(DATA_ROOT, 'train_spectrograms')\n\n# ── reproducibility ───────────────────────────────────────────────────────────\nrandom.seed(cfg.seed)\nnp.random.seed(cfg.seed)\ntorch.manual_seed(cfg.seed)\ntorch.cuda.manual_seed_all(cfg.seed)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark     = False\n\nDEVICE  = torch.device(cfg.device)\nUSE_AMP = DEVICE.type == 'cuda'\n\nprint(cfg)\nprint(f\"Device: {DEVICE} | AMP: {USE_AMP}\")\n\n# NOTE: if running on Kaggle and this path doesn't exist, run !ls /kaggle/input\n# and update DATA_ROOT to match how the competition data was attached.\nprint(f\"Competition data path exists: {os.path.exists(DATA_ROOT)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell-351296","cell_type":"code","source":"# ============ Data loading (workshop subset) ============\n# The raw HMS train.csv carries everything SpectrogramDataset needs (spectrogram_id,\n# spectrogram_label_offset_seconds, vote columns). workshop_sample_ids.csv only carries\n# ID + split, so we merge the two to reconstruct full rows for the 600/120 workshop subset.\n\nraw_df     = pd.read_csv(RAW_TRAIN_PATH)\nsample_ids = pd.read_csv(SAMPLE_IDS_PATH)\n\nmerged = raw_df.merge(\n    sample_ids[[\"eeg_id\", \"eeg_sub_id\", \"split\"]],\n    on=[\"eeg_id\", \"eeg_sub_id\"],\n    how=\"inner\",\n)\ntrain_df = merged[merged[\"split\"] == \"train\"].reset_index(drop=True)\nval_df   = merged[merged[\"split\"] == \"val\"].reset_index(drop=True)\n\nprint(f'Train: {len(train_df):,} rows')\nprint(f'Val   : {len(val_df):,} rows')\nprint(train_df['expert_consensus'].value_counts())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell-92f0fa","cell_type":"code","source":"# ============ Spectrogram preprocessing ============\n# Kaggle's precomputed spectrograms already have 4 bipolar chains (LL, RL, LP, RP),\n# each with 100 frequency bins, stacked into 400 rows per parquet file.\n# We convert each {spec_id}.parquet → {spec_id}.npy once, then cache it.\n\nSEED = Config2D.seed\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\n\ndef preprocess_one_spectrogram(spec_id: int, spectrogram_dir: str, cache_dir: str) -> None:\n    dst = os.path.join(cache_dir, f\"{spec_id}.npy\")\n    if os.path.exists(dst):\n        return\n    src = os.path.join(spectrogram_dir, f\"{spec_id}.parquet\")\n    df  = pd.read_parquet(src)\n    df  = df.fillna(0)\n    if \"time\" in df.columns:\n        df = df.drop(columns=[\"time\"])\n    # (total_time, 400) → transpose → (400, total_time)\n    arr = df.to_numpy(dtype=np.float32).T\n    np.save(dst, arr)\n\n\nos.makedirs(cfg.spec_cache_dir, exist_ok=True)\n\n# Workshop mode only needs the spectrograms actually used by the 720-row subset —\n# far fewer than the full ~11,000 files, which keeps this step fast (seconds, not minutes).\nneeded_ids = set(train_df['spectrogram_id']).union(val_df['spectrogram_id'])\nfor spec_id in needed_ids:\n    preprocess_one_spectrogram(int(spec_id), cfg.spectrogram_dir, cfg.spec_cache_dir)\n\nprint(f'Cache ready: {len(needed_ids)} spectrograms for workshop subset')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell-d4f29d","cell_type":"code","source":"# ============ Dataset ============\nVOTE_COLS = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\",\n             \"lrda_vote\",    \"grda_vote\", \"other_vote\"]\n\n\nclass SpectrogramDataset(Dataset):\n    \"\"\"\n    Loads pre-cached (400, total_time) .npy spectrograms, extracts a 300-column\n    (600-second) window around the labeled offset, normalizes it, and reshapes\n    the 400 stacked rows back into 4 separate bipolar chains: (4, 100, 300).\n    \"\"\"\n\n    def __init__(self, metadata_df: pd.DataFrame, spec_cache_dir: str):\n        self.meta           = metadata_df.reset_index(drop=True)\n        self.spec_cache_dir = spec_cache_dir\n\n        votes    = self.meta[VOTE_COLS].to_numpy(dtype=np.float32)\n        row_sums = votes.sum(axis=1, keepdims=True)\n        row_sums = np.where(row_sums == 0, 1.0, row_sums)\n        self.soft_labels = votes / row_sums          # (N, 6)\n        self.hard_labels = self.soft_labels.argmax(1)  # (N,)\n\n    def __len__(self) -> int:\n        return len(self.meta)\n\n    def __getitem__(self, idx: int) -> dict:\n        row     = self.meta.iloc[idx]\n        spec_id = int(row[\"spectrogram_id\"])\n        offset  = float(row[\"spectrogram_label_offset_seconds\"])\n\n        spec = np.load(os.path.join(self.spec_cache_dir, f\"{spec_id}.npy\"))\n\n        # extract 300-column window; 2-second time resolution → col = offset // 2\n        col_start = int(offset // 2)\n        window    = spec[:, col_start:col_start + 300]   # (400, ≤300)\n\n        # right-pad to exactly 300 columns if the window runs off the end\n        if window.shape[1] < 300:\n            pad    = 300 - window.shape[1]\n            window = np.pad(window, ((0, 0), (0, pad)), mode=\"constant\")\n\n        # normalize: clip → log → z-score\n        window = np.clip(window, np.exp(-4), np.exp(8))\n        window = np.log(window)\n        mu     = window.mean()\n        sigma  = window.std()\n        window = (window - mu) / (sigma + 1e-6)\n\n        # split 400 stacked rows back into 4 chains: (400, 300) → (4, 100, 300)\n        window = window.reshape(4, 100, 300)\n\n        img   = torch.from_numpy(window).float()\n        soft  = torch.from_numpy(self.soft_labels[idx])\n        label = int(self.hard_labels[idx])\n\n        return {\"image\": img, \"soft_label\": soft, \"label\": label}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell-da12b8","cell_type":"code","source":"# ============ Model ============\nclass EfficientNetEEG(nn.Module):\n    \"\"\"\n    timm EfficientNet-B0 backbone adapted for 4-channel spectrogram input\n    (instead of the usual 3-channel RGB), with a linear head for 6 classes.\n    Pretrained on ImageNet — note the domain gap: natural images vs EEG spectrograms.\n    \"\"\"\n\n    def __init__(self, backbone=\"efficientnet_b0\", num_classes=6,\n                 pretrained=True, drop_rate=0.3):\n        super().__init__()\n        self.encoder = timm.create_model(\n            backbone,\n            pretrained=pretrained,\n            in_chans=4,       # 4 bipolar chains instead of RGB\n            num_classes=0,    # remove timm's default head; add our own\n            drop_rate=drop_rate,\n        )\n        n_features = self.encoder.num_features\n        self.head = nn.Linear(n_features, num_classes)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        features = self.encoder(x)\n        return self.head(features)\n\n\nmodel = EfficientNetEEG(\n    backbone    = cfg.backbone,\n    num_classes = cfg.num_classes,\n    pretrained  = cfg.pretrained,\n    drop_rate   = cfg.drop_rate,\n).to(DEVICE)\n\ntotal_params     = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Backbone    : {cfg.backbone}\")\nprint(f\"Parameters  — total: {total_params:,} | trainable: {trainable_params:,}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell-ee6dec","cell_type":"code","source":"# ============ Single-stage training ============\nGRAD_CLIP  = 1.0\nNUM_EPOCHS = 5   # ~5-10 min on Kaggle T4 for this subset size\nLR         = cfg.lr\n\ncriterion     = nn.KLDivLoss(reduction='batchmean')\nval_criterion = nn.CrossEntropyLoss()\nscaler        = torch.amp.GradScaler('cuda', enabled=USE_AMP)\n\nhistory = []\n\ntrain_ds     = SpectrogramDataset(train_df, cfg.spec_cache_dir)\ntrain_loader = DataLoader(train_ds, batch_size=cfg.batch_size, shuffle=True,\n                          num_workers=cfg.num_workers, pin_memory=True)\nval_ds       = SpectrogramDataset(val_df, cfg.spec_cache_dir)\nval_loader   = DataLoader(val_ds, batch_size=cfg.batch_size, shuffle=False,\n                          num_workers=cfg.num_workers, pin_memory=True)\nprint(f'Train: {len(train_ds):,} | Val: {len(val_ds):,} | Epochs: {NUM_EPOCHS}')\n\n\ndef validate_2d(model, loader):\n    model.eval()\n    val_ce, val_kl = 0.0, 0.0\n    all_preds, all_labels = [], []\n    with torch.no_grad():\n        for batch in loader:\n            x    = batch['image'].to(DEVICE)\n            y    = batch['label'].to(DEVICE)\n            soft = batch['soft_label'].to(DEVICE)\n            with torch.amp.autocast('cuda', enabled=USE_AMP):\n                logits = model(x)\n            val_ce += val_criterion(logits, y).item()\n            val_kl += criterion(F.log_softmax(logits, dim=1), soft).item()\n            all_preds .extend(logits.argmax(1).cpu().tolist())\n            all_labels.extend(y.cpu().tolist())\n    macro_f1  = f1_score(all_labels, all_preds, average='macro', zero_division=0)\n    per_class = f1_score(all_labels, all_preds, average=None,    zero_division=0)\n    return (val_kl / len(loader), val_ce / len(loader), macro_f1, per_class)\n\n\noptimizer = AdamW(model.parameters(), lr=LR, weight_decay=cfg.weight_decay)\nscheduler = CosineAnnealingLR(optimizer, T_max=NUM_EPOCHS)\n# NOTE: \"best\" is defined by val KL (the competition metric), matching the XGBoost\n# and 1D CNN notebooks, so all three models are selected the same way.\nbest_kl, best_f1, best_epoch, wait = float('inf'), 0.0, 0, 0\nbest_state = None\n\nfor epoch in range(1, NUM_EPOCHS + 1):\n    model.train()\n    t0 = time.time()\n    train_loss = 0.0\n    for batch in train_loader:\n        x    = batch['image'].to(DEVICE)\n        soft = batch['soft_label'].to(DEVICE)\n        optimizer.zero_grad()\n        with torch.amp.autocast('cuda', enabled=USE_AMP):\n            logits = model(x)\n            loss   = criterion(F.log_softmax(logits, dim=1), soft)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        scaler.step(optimizer)\n        scaler.update()\n        train_loss += loss.item()\n    scheduler.step()\n\n    val_kl, val_ce, macro_f1, _ = validate_2d(model, val_loader)\n    avg_train = train_loss / len(train_loader)\n    elapsed   = time.time() - t0\n\n    history.append({\n        'epoch': epoch, 'train_kl': avg_train,\n        'val_kl': val_kl, 'val_ce': val_ce, 'macro_f1': macro_f1,\n    })\n    print(f'Epoch {epoch:03d} | train_kl {avg_train:.4f} | '\n          f'val_kl {val_kl:.4f} | val_ce {val_ce:.4f} | '\n          f'macro_f1 {macro_f1:.4f} | {elapsed:.0f}s')\n\n    if val_kl < best_kl:\n        best_kl, best_f1, best_epoch, wait = val_kl, macro_f1, epoch, 0\n        best_state = {k: v.cpu().clone() for k, v in model.state_dict().items()}\n        torch.save(model.state_dict(), 'best_2d_workshop.pt')\n        print(f'  \\u2713 saved best_2d_workshop.pt (val_kl={best_kl:.4f})')\n    else:\n        wait += 1\n        if wait >= cfg.patience:\n            print(f'Early stopping at epoch {epoch}')\n            break\n\nprint(f'\\nTraining complete \\u2014 best val_kl={best_kl:.4f} (macro_f1={best_f1:.4f}) at epoch {best_epoch}')\nif best_state is not None:\n    model.load_state_dict(best_state)\n\n_, _, _, best_per_class = validate_2d(model, val_loader)  # per-class F1 at the KL-best checkpoint\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"cell-a523f0","cell_type":"code","source":"# ============ Final report (standard format) ============\nCLASS_NAMES = [\"Seizure\", \"LPD\", \"GPD\", \"LRDA\", \"GRDA\", \"Other\"]\n\nepochs_   = [h['epoch']    for h in history]\ntrain_kl_ = [h['train_kl'] for h in history]\nval_kl_   = [h['val_kl']   for h in history]\nval_ce_   = [h['val_ce']   for h in history]\nmacro_f1_ = [h['macro_f1'] for h in history]\n\nprint(f\"Val macro F1     : {best_f1:.4f}\")\nprint(f\"Val KL divergence: {best_kl:.4f}\")\nprint(f\"(random-guess baseline for 6 balanced classes: macro F1 \\u2248 0.167)\")\nprint(\"\\nPer-class F1 (best epoch):\")\nfor name, f in zip(CLASS_NAMES, best_per_class):\n    print(f\"  {name:<10} {f:.4f}\")\n\nfig, axes = plt.subplots(2, 2, figsize=(13, 9))\n\naxes[0, 0].plot(epochs_, train_kl_, label='train KL (loss)')\naxes[0, 0].plot(epochs_, val_kl_,   label='val KL (loss)')\naxes[0, 0].axvline(best_epoch, color='red', linestyle='--', label=f'best={best_epoch}')\naxes[0, 0].set_xlabel('Epoch'); axes[0, 0].set_ylabel('KL Divergence (loss)')\naxes[0, 0].set_title('Train/Val KL loss'); axes[0, 0].legend(fontsize=8)\n\naxes[0, 1].plot(epochs_, val_ce_, color='orange')\naxes[0, 1].axvline(best_epoch, color='red', linestyle='--', label=f'best={best_epoch}')\naxes[0, 1].set_xlabel('Epoch'); axes[0, 1].set_ylabel('Cross-Entropy Loss')\naxes[0, 1].set_title('Val CE (monitoring)'); axes[0, 1].legend(fontsize=8)\n\naxes[1, 0].plot(epochs_, macro_f1_, color='seagreen')\naxes[1, 0].axhline(1/6, color='gray', linestyle='--', linewidth=1, label='random guess')\naxes[1, 0].axvline(best_epoch, color='red', linestyle='--', label=f'best={best_epoch}')\naxes[1, 0].set_xlabel('Epoch'); axes[1, 0].set_ylabel('Val Macro F1')\naxes[1, 0].set_title('Val Macro F1'); axes[1, 0].legend(fontsize=8)\n\naxes[1, 1].bar(CLASS_NAMES, best_per_class, color='steelblue')\naxes[1, 1].axhline(best_f1, color='red', linestyle='--', label=f'macro F1 = {best_f1:.3f}')\naxes[1, 1].set_ylim(0, 1); axes[1, 1].set_ylabel('F1')\naxes[1, 1].set_title(f'Per-class F1 (val, epoch {best_epoch})'); axes[1, 1].legend(fontsize=8)\n\nplt.tight_layout()\nplt.savefig('training_curves_2d_workshop.png', dpi=150)\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"f91d0280","cell_type":"markdown","source":"## Tweak & Compare (~20 min)\n\nPick 1-2 changes below, rerun the cell, and log your result in the shared sheet.\n\n| Knob | Default | Try |\n|---|---|---|\n| learning rate | 1e-3 | 3e-4 / 3e-3 |\n| `drop_rate` | 0.3 | 0.1 / 0.5 |\n| time window (`img_width`) | 300 (600 s) | 150 (300 s) / 450 (900 s) |\n\n**Log your run:** What you changed | Val F1 | Val KL | one-line observation\n","metadata":{}},{"id":"d35bee87","cell_type":"code","source":"# ============ Tweak & Compare ============\nTWEAK_LR         = 1e-3   # <- change me (try 3e-4 or 3e-3)\nTWEAK_DROP_RATE  = 0.3    # <- change me (try 0.1 or 0.5)\nTWEAK_IMG_WIDTH  = 300    # <- change me (try 150 or 450)\n\n# re-seed locally so this cell gives the same result every time it's rerun on its own,\n# regardless of how many earlier cells (and how much of the global RNG stream) ran first\nSEED = Config2D.seed\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\n\n\nclass SpectrogramDatasetTweak(SpectrogramDataset):\n    \"\"\"Same as SpectrogramDataset, but with a configurable time-window width\n    (instead of the fixed 300 columns / 600 seconds).\"\"\"\n\n    def __init__(self, metadata_df, spec_cache_dir, width=300):\n        self.width = width\n        super().__init__(metadata_df, spec_cache_dir)\n\n    def __getitem__(self, idx):\n        row     = self.meta.iloc[idx]\n        spec_id = int(row[\"spectrogram_id\"])\n        offset  = float(row[\"spectrogram_label_offset_seconds\"])\n\n        spec = np.load(os.path.join(self.spec_cache_dir, f\"{spec_id}.npy\"))\n        col_start = int(offset // 2)\n        window    = spec[:, col_start:col_start + self.width]\n\n        if window.shape[1] < self.width:\n            pad = self.width - window.shape[1]\n            window = np.pad(window, ((0, 0), (0, pad)), mode=\"constant\")\n\n        # normalize: clip -> log -> z-score (same as SpectrogramDataset above —\n        # missing this step is what caused the earlier Val KL: inf / F1: 0.0000 bug:\n        # raw un-logged spectrogram power values are unbounded and blow up the\n        # pretrained EfficientNet's activations within a few batches)\n        window = np.clip(window, np.exp(-4), np.exp(8))\n        window = np.log(window)\n        mu     = window.mean()\n        sigma  = window.std()\n        window = (window - mu) / (sigma + 1e-6)\n\n        image = window.reshape(4, cfg.img_height, self.width).astype(np.float32)\n        image = np.nan_to_num(image, nan=0.0, posinf=0.0, neginf=0.0)\n\n        return {\n            \"image\":      torch.from_numpy(image),\n            \"label\":      int(self.hard_labels[idx]),\n            \"soft_label\": torch.from_numpy(self.soft_labels[idx]),\n        }\n\n\ntweak_train_ds = SpectrogramDatasetTweak(train_df, cfg.spec_cache_dir, width=TWEAK_IMG_WIDTH)\ntweak_val_ds   = SpectrogramDatasetTweak(val_df,   cfg.spec_cache_dir, width=TWEAK_IMG_WIDTH)\ntweak_train_loader = DataLoader(tweak_train_ds, batch_size=cfg.batch_size, shuffle=True,\n                                 num_workers=cfg.num_workers, pin_memory=True)\ntweak_val_loader   = DataLoader(tweak_val_ds,   batch_size=cfg.batch_size, shuffle=False,\n                                 num_workers=cfg.num_workers, pin_memory=True)\n\ntweak_model = EfficientNetEEG(\n    backbone=cfg.backbone, num_classes=cfg.num_classes,\n    pretrained=cfg.pretrained, drop_rate=TWEAK_DROP_RATE,\n).to(DEVICE)\ntweak_optimizer = AdamW(tweak_model.parameters(), lr=TWEAK_LR, weight_decay=cfg.weight_decay)\ntweak_scheduler = CosineAnnealingLR(tweak_optimizer, T_max=NUM_EPOCHS)\ntweak_scaler    = torch.amp.GradScaler('cuda', enabled=USE_AMP)\n\nt_best_kl, t_best_f1, t_wait = float('inf'), 0.0, 0\nfor epoch in range(1, NUM_EPOCHS + 1):\n    tweak_model.train()\n    for batch in tweak_train_loader:\n        x    = batch['image'].to(DEVICE)\n        soft = batch['soft_label'].to(DEVICE)\n        tweak_optimizer.zero_grad()\n        with torch.amp.autocast('cuda', enabled=USE_AMP):\n            logits = tweak_model(x)\n            loss   = criterion(F.log_softmax(logits, dim=1), soft)\n        tweak_scaler.scale(loss).backward()\n        tweak_scaler.unscale_(tweak_optimizer)\n        nn.utils.clip_grad_norm_(tweak_model.parameters(), GRAD_CLIP)\n        tweak_scaler.step(tweak_optimizer)\n        tweak_scaler.update()\n    tweak_scheduler.step()\n\n    t_val_kl, _, t_macro_f1, _ = validate_2d(tweak_model, tweak_val_loader)\n    if t_val_kl < t_best_kl:\n        t_best_kl, t_best_f1, t_wait = t_val_kl, t_macro_f1, 0\n    else:\n        t_wait += 1\n        if t_wait >= cfg.patience:\n            break\n\nprint(f\"lr={TWEAK_LR} | drop_rate={TWEAK_DROP_RATE} | img_width={TWEAK_IMG_WIDTH}\")\nprint(f\"Val macro F1 : {t_best_f1:.4f}\")\nprint(f\"Val KL       : {t_best_kl:.4f}\")\nprint(\"-> log this row in the shared sheet: what you changed / F1 / KL / one-line observation\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"3c204052","cell_type":"markdown","source":"### Done early? Extend (~10 min)\n\n**What this does:** ensembles this 2D CNN's predictions with the 1D CNN's (domain-informed)\npredictions, and checks whether the combination beats either model alone. This is one of the\nmost consistent findings across the top-10 Kaggle solutions review earlier — and here you get\nto see it happen with two models built from genuinely different information sources\n(time-domain waveform vs. frequency-domain spectrogram of the *same* labeled event).\n\nThis step involves a bit of cross-notebook plumbing (matching two independently-loaded\nvalidation sets by ID, reusing ground-truth labels from this notebook's own val set since the\n1D CNN's exported file doesn't carry them) that isn't a great fit for a quick AI-prompt exercise\nat this point in a long day — so unlike the earlier Tweak & Extend steps, just **run the cell\nbelow directly** and read through what it's doing.\n\nA pre-computed `cnn1d_domain_informed_val_probs.csv` (from the 1D CNN notebook's export cell)\nneeds to be attached to this notebook as a dataset for this to work. Update the path below to\nmatch — Kaggle attaches datasets under `/kaggle/input/<owner>/<dataset-slug>/`; check the file\nbrowser panel on the right if the path doesn't match.\n","metadata":{}},{"id":"a8fbaa53","cell_type":"code","source":"# ============ Ensemble: 2D CNN + 1D CNN ============\nCNN1D_PROBS_PATH = \"/kaggle/input/datasets/xiaosufrankhu/midas-summer-academy-wk3-eeg/cnn1d_domain_informed_val_probs.csv\"\ncnn1d_probs_df = pd.read_csv(CNN1D_PROBS_PATH)\n\n# ---- get this 2D CNN's own val-set probabilities (using the KL-best checkpoint,\n#      already loaded back into `model` at the end of the Single-stage training cell) ----\nmodel.eval()\ncnn2d_ids, cnn2d_probs, cnn2d_soft = [], [], []\nwith torch.no_grad():\n    for batch in val_loader:\n        x = batch['image'].to(DEVICE)\n        with torch.amp.autocast('cuda', enabled=USE_AMP):\n            logits = model(x)\n        cnn2d_probs.append(F.softmax(logits, dim=1).cpu().numpy())\n        cnn2d_soft.append(batch['soft_label'].numpy())\ncnn2d_probs = np.concatenate(cnn2d_probs, axis=0)   # (120, 6)\ncnn2d_soft  = np.concatenate(cnn2d_soft,  axis=0)   # (120, 6), true vote distribution\n\nprob_cols = [f\"prob_{c.lower()}\" for c in CLASS_NAMES]\ncnn2d_ids_df = val_df[[\"eeg_id\", \"eeg_sub_id\"]].reset_index(drop=True)\ncnn2d_df = pd.concat([cnn2d_ids_df, pd.DataFrame(cnn2d_probs, columns=prob_cols)], axis=1)\n\n# ---- merge on ID rather than assuming row order matches — the two notebooks\n#      load/order the validation set independently ----\nmerged = cnn2d_df.merge(cnn1d_probs_df, on=[\"eeg_id\", \"eeg_sub_id\"], suffixes=(\"_2d\", \"_1d\"))\nassert len(merged) == len(cnn2d_df), \"some validation rows didn't find a 1D CNN match — check the CSV path/upload\"\n\nprobs_2d = merged[[f\"{c}_2d\" for c in prob_cols]].to_numpy()\nprobs_1d = merged[[f\"{c}_1d\" for c in prob_cols]].to_numpy()\nprobs_ensemble = (probs_2d + probs_1d) / 2\n\n# re-derive the true soft-label distribution in the merged row order\nsoft_lookup = dict(zip(zip(cnn2d_ids_df[\"eeg_id\"], cnn2d_ids_df[\"eeg_sub_id\"]),\n                        list(cnn2d_soft)))\nsoft_merged = np.stack([soft_lookup[(r.eeg_id, r.eeg_sub_id)] for r in merged.itertuples()])\nhard_merged = soft_merged.argmax(axis=1)\n\n\ndef _f1_kl(probs, hard_labels, soft_labels):\n    f1 = f1_score(hard_labels, probs.argmax(axis=1), average=\"macro\", zero_division=0)\n    kl = (soft_labels * np.log(np.clip(soft_labels, 1e-7, 1) / np.clip(probs, 1e-7, 1))).sum(axis=1).mean()\n    return f1, kl\n\n\nf1_2d, kl_2d = _f1_kl(probs_2d, hard_merged, soft_merged)\nf1_1d, kl_1d = _f1_kl(probs_1d, hard_merged, soft_merged)\nf1_ens, kl_ens = _f1_kl(probs_ensemble, hard_merged, soft_merged)\n\nprint(f\"{'Model':<20}{'Val F1':>10}{'Val KL':>10}\")\nprint(f\"{'-'*40}\")\nprint(f\"{'2D CNN alone':<20}{f1_2d:>10.4f}{kl_2d:>10.4f}\")\nprint(f\"{'1D CNN alone':<20}{f1_1d:>10.4f}{kl_1d:>10.4f}\")\nprint(f\"{'Ensemble (avg)':<20}{f1_ens:>10.4f}{kl_ens:>10.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"a4a408d1-9013-444d-8688-cc7d864f9848","cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}