{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Cell 1: Setup and imports\nimport os\nimport gc\nimport json\nimport time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nfrom tqdm.auto import tqdm\n\n# Paths\nRAW_DIR = Path('/kaggle/input/competitions/bengaliai-cv19')\nWORK_DIR = Path('/kaggle/working')\nCACHE_DIR = WORK_DIR / 'cache'\nFIG_DIR = WORK_DIR / 'figures'\nCACHE_DIR.mkdir(parents=True, exist_ok=True)\nFIG_DIR.mkdir(parents=True, exist_ok=True)\n\n# Constants — dataset's fixed dimensions\nRAW_H, RAW_W = 137, 236\nTARGET_SIZE = 128\nN_TOTAL_EXPECTED = 200_840\nN_ROOT_CLASSES = 168\nN_VOWEL_CLASSES = 11\nN_CONSONANT_CLASSES = 7   # original scheme; we're not using class_map_corrected\n\nprint(\"Environment ready.\")\nprint(f\"  Raw data dir: {RAW_DIR}\")\nprint(f\"  Working dir:  {WORK_DIR}\")\nprint(f\"  Cache dir:    {CACHE_DIR}\")\nprint(f\"  Figure dir:   {FIG_DIR}\")\n\n# Confirm the dataset is attached\nassert RAW_DIR.exists(), (\n    f\"{RAW_DIR} not found. Add the 'Bengali.AI Handwritten Grapheme Classification' \"\n    f\"dataset to your notebook (Add Data → search 'bengaliai-cv19').\"\n)\n\nraw_files = sorted(RAW_DIR.glob('*'))\nprint(f\"\\nFiles in raw dir ({len(raw_files)}):\")\nfor f in raw_files:\n    print(f\"  {f.name:40s}  ({f.stat().st_size / 1e6:7.1f} MB)\")\n\n# Verify all 4 training parquets are present\ntrain_parquets = sorted(RAW_DIR.glob('train_image_data_*.parquet'))\nassert len(train_parquets) == 4, (\n    f\"Expected 4 train parquets, found {len(train_parquets)}. \"\n    f\"Dataset may be incomplete.\"\n)\nprint(f\"\\n✓ All 4 training parquets found\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:07:01.793183Z","iopub.execute_input":"2026-09-19T06:07:01.794085Z","iopub.status.idle":"2026-09-19T06:07:01.813732Z","shell.execute_reply.started":"2026-09-19T06:07:01.794043Z","shell.execute_reply":"2026-09-19T06:07:01.813094Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 2: Load labels, verify, and produce class-distribution figure\ndf = pd.read_csv(RAW_DIR / 'train.csv')\nprint(f\"Loaded train.csv: {len(df):,} rows\")\nprint(f\"Columns: {list(df.columns)}\")\nprint(df.head())\n\n# ----- Sanity assertions -----\nassert len(df) == N_TOTAL_EXPECTED, f\"Expected {N_TOTAL_EXPECTED} rows, got {len(df)}\"\nassert df['grapheme_root'].nunique() == N_ROOT_CLASSES, \\\n    f\"Expected {N_ROOT_CLASSES} roots, got {df['grapheme_root'].nunique()}\"\nassert df['vowel_diacritic'].nunique() == N_VOWEL_CLASSES, \\\n    f\"Expected {N_VOWEL_CLASSES} vowels, got {df['vowel_diacritic'].nunique()}\"\nassert df['consonant_diacritic'].nunique() == N_CONSONANT_CLASSES, \\\n    f\"Expected {N_CONSONANT_CLASSES} consonants, got {df['consonant_diacritic'].nunique()}\"\n\n# Verify class-id ranges are 0..N-1 (no gaps)\nfor col, n in [('grapheme_root', N_ROOT_CLASSES),\n               ('vowel_diacritic', N_VOWEL_CLASSES),\n               ('consonant_diacritic', N_CONSONANT_CLASSES)]:\n    vals = set(df[col].unique())\n    expected = set(range(n))\n    assert vals == expected, f\"{col} has unexpected values: {vals - expected} or missing {expected - vals}\"\n\n# Verify image_ids are unique\nassert df['image_id'].is_unique, \"Duplicate image_ids in train.csv!\"\n\nn_unique_wholes = df['grapheme'].nunique()\nprint(f\"\\n✓ All assertions pass\")\nprint(f\"  Grapheme roots:      {df['grapheme_root'].nunique()}\")\nprint(f\"  Vowel diacritics:    {df['vowel_diacritic'].nunique()}\")\nprint(f\"  Consonant diacritics:{df['consonant_diacritic'].nunique()}\")\nprint(f\"  Unique wholes:       {n_unique_wholes}\")\n\n# ----- Class-frequency figure -----\nfig, axes = plt.subplots(1, 3, figsize=(18, 4))\n\nfor ax, col, title, color in [\n    (axes[0], 'grapheme_root',      f'Grapheme root ({N_ROOT_CLASSES} classes)',        '#C0392B'),\n    (axes[1], 'vowel_diacritic',    f'Vowel diacritic ({N_VOWEL_CLASSES} classes)',     '#1F6F6B'),\n    (axes[2], 'consonant_diacritic',f'Consonant diacritic ({N_CONSONANT_CLASSES} classes)', '#1F6F6B'),\n]:\n    counts = df[col].value_counts().sort_values(ascending=False).values\n    ax.bar(range(len(counts)), counts, color=color)\n    ax.set_yscale('log')\n    ax.set_title(title, fontsize=13)\n    ax.set_xlabel('Class rank (most→least frequent)')\n    ax.set_ylabel('Count (log scale)')\n    ax.grid(axis='y', alpha=0.3)\n    ratio = counts.max() / max(counts.min(), 1)\n    ax.text(0.98, 0.95, f'Max/min ratio: {ratio:,.0f}×',\n            transform=ax.transAxes, ha='right', va='top', fontsize=11,\n            bbox=dict(facecolor='white', edgecolor='gray', alpha=0.9))\n\nplt.suptitle('Class-frequency distributions — why we need macro-recall, not accuracy',\n             fontsize=14, y=1.02)\nplt.tight_layout()\nplt.savefig(FIG_DIR / 'eda_class_distribution.png', dpi=140, bbox_inches='tight')\nplt.show()\n\nprint(f\"\\n✓ Figure saved to {FIG_DIR / 'eda_class_distribution.png'}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:07:21.53837Z","iopub.execute_input":"2026-09-19T06:07:21.539191Z","iopub.status.idle":"2026-09-19T06:07:23.491768Z","shell.execute_reply.started":"2026-09-19T06:07:21.539161Z","shell.execute_reply":"2026-09-19T06:07:23.491008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 3: Preprocessing function and single-image sanity test\ndef preprocess_image(flat_pixels, target_size=TARGET_SIZE, ink_threshold=200):\n    \"\"\"Turn a raw parquet row into a cached uint8 image.\n\n    Steps:\n    1. Reshape flat pixel array to (137, 236)\n    2. Invert: raw is dark ink on light paper → bright ink on dark background\n       (matches typical CNN-friendly convention: signal = high values)\n    3. Threshold-crop to the ink bounding box (removes wasted background)\n    4. Pad to square (preserves aspect ratio, centered by construction)\n    5. Resize to target_size × target_size with area interpolation\n\n    Returns: uint8 array (target_size, target_size)\n    \"\"\"\n    # 1. Reshape\n    img = flat_pixels.reshape(RAW_H, RAW_W).astype(np.uint8)\n\n    # 2. Invert (raw: ink is dark ~0, paper is light ~255)\n    img = 255 - img\n\n    # 3. Ink bounding-box crop\n    # After inversion, ink pixels are bright. Anything with intensity > (255 - ink_threshold)\n    # is considered ink; the rest is background.\n    ink_mask = img > (255 - ink_threshold)\n    if ink_mask.any():\n        rows = np.any(ink_mask, axis=1)\n        cols = np.any(ink_mask, axis=0)\n        rmin, rmax = np.where(rows)[0][[0, -1]]\n        cmin, cmax = np.where(cols)[0][[0, -1]]\n        margin = 3\n        rmin = max(0, rmin - margin)\n        rmax = min(RAW_H - 1, rmax + margin)\n        cmin = max(0, cmin - margin)\n        cmax = min(RAW_W - 1, cmax + margin)\n        img = img[rmin:rmax + 1, cmin:cmax + 1]\n    # If image is fully blank (shouldn't happen but safe), fall through with the raw image\n\n    # 4. Pad to square (put the shorter dim into the longer dim's frame)\n    h, w = img.shape\n    side = max(h, w)\n    pad_h = (side - h) // 2\n    pad_w = (side - w) // 2\n    padded = np.zeros((side, side), dtype=np.uint8)\n    padded[pad_h:pad_h + h, pad_w:pad_w + w] = img\n\n    # 5. Resize\n    resized = cv2.resize(padded, (target_size, target_size), interpolation=cv2.INTER_AREA)\n    return resized\n\n\n# ----- Single-image test on the first parquet's first row -----\n# Read only the first row (pyarrow can't do row-slicing on parquet, but reading\n# the whole file just to preview one row is fine here — it's a one-time check)\nsample_pq = pd.read_parquet(RAW_DIR / 'train_image_data_0.parquet').iloc[0]\nsample_id = sample_pq['image_id']\n\n# All columns except image_id are pixel columns, named \"0\", \"1\", ..., \"32331\"\npixel_cols = [c for c in sample_pq.index if c != 'image_id']\nassert len(pixel_cols) == RAW_H * RAW_W, \\\n    f\"Expected {RAW_H * RAW_W} pixel columns, got {len(pixel_cols)}\"\n\nsample_pixels = sample_pq[pixel_cols].values.astype(np.uint8)\nsample_processed = preprocess_image(sample_pixels)\n\nprint(f\"✓ Preprocessing function works\")\nprint(f\"  Sample image_id:   {sample_id}\")\nprint(f\"  Raw pixel count:   {sample_pixels.shape[0]} → reshape to ({RAW_H}, {RAW_W})\")\nprint(f\"  Processed shape:   {sample_processed.shape}\")\nprint(f\"  Processed dtype:   {sample_processed.dtype}\")\nprint(f\"  Value range:       [{sample_processed.min()}, {sample_processed.max()}]\")\n\n# Show raw vs. processed\nfig, axes = plt.subplots(1, 2, figsize=(11, 4.5))\naxes[0].imshow(sample_pixels.reshape(RAW_H, RAW_W), cmap='gray')\naxes[0].set_title(f'Raw (137×236)\\nDark ink on light paper')\naxes[0].axis('off')\naxes[1].imshow(sample_processed, cmap='gray')\naxes[1].set_title(f'Processed (128×128)\\nBright ink on dark, cropped & centered')\naxes[1].axis('off')\nplt.tight_layout()\nplt.show()\n\n# Free memory before Cell 4\ndel sample_pq, sample_pixels\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:07:36.97415Z","iopub.execute_input":"2026-09-19T06:07:36.974927Z","iopub.status.idle":"2026-09-19T06:07:51.223163Z","shell.execute_reply.started":"2026-09-19T06:07:36.974897Z","shell.execute_reply":"2026-09-19T06:07:51.222527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 4: Preprocess all images and write to memmap\nCACHE_IMAGES_PATH = CACHE_DIR / 'images_128.npy'\nCACHE_LABELS_PATH = CACHE_DIR / 'labels.parquet'\n\nif CACHE_IMAGES_PATH.exists() and CACHE_LABELS_PATH.exists():\n    print(f\"✓ Cache already exists at {CACHE_IMAGES_PATH}\")\n    print(f\"  Size: {CACHE_IMAGES_PATH.stat().st_size / 1e9:.2f} GB\")\n    print(f\"  Delete it manually if you want to regenerate.\")\nelse:\n    print(f\"Building cache at {CACHE_IMAGES_PATH}\")\n    expected_size_gb = N_TOTAL_EXPECTED * TARGET_SIZE * TARGET_SIZE / 1e9\n    print(f\"Expected size: {N_TOTAL_EXPECTED} × {TARGET_SIZE} × {TARGET_SIZE} \"\n          f\"= {expected_size_gb:.2f} GB\\n\")\n\n    # Pre-allocate the full memmap on disk (avoids OOM)\n    images = np.lib.format.open_memmap(\n        CACHE_IMAGES_PATH, mode='w+', dtype=np.uint8,\n        shape=(N_TOTAL_EXPECTED, TARGET_SIZE, TARGET_SIZE)\n    )\n\n    # Canonical row order: match train.csv\n    id_to_idx = {img_id: i for i, img_id in enumerate(df['image_id'].values)}\n\n    t0 = time.time()\n    processed_count = 0\n\n    for pq_idx in range(4):\n        pq_path = RAW_DIR / f'train_image_data_{pq_idx}.parquet'\n        print(f\"[{pq_idx + 1}/4] Reading {pq_path.name}...\")\n        pq_df = pd.read_parquet(pq_path)\n\n        # Extract image_ids and pixel matrix in bulk (this is the speedup)\n        ids = pq_df['image_id'].values\n        pixel_cols = [c for c in pq_df.columns if c != 'image_id']\n        # (n_rows_in_this_parquet, 32332) as uint8\n        pixel_matrix = pq_df[pixel_cols].values.astype(np.uint8)\n\n        # Free the DataFrame — pixel_matrix is what we need\n        del pq_df\n        gc.collect()\n\n        print(f\"        Processing {len(ids):,} images...\")\n        for i in tqdm(range(len(ids)), desc=f'parquet {pq_idx}'):\n            img_id = ids[i]\n            idx = id_to_idx[img_id]\n            images[idx] = preprocess_image(pixel_matrix[i])\n            processed_count += 1\n\n        # Free this parquet's pixels before loading the next\n        del ids, pixel_matrix\n        gc.collect()\n\n    # Flush memmap to disk\n    images.flush()\n    del images\n    gc.collect()\n\n    elapsed = time.time() - t0\n    print(f\"\\n✓ Processed {processed_count:,} images in {elapsed / 60:.1f} min\")\n    print(f\"  Rate: {processed_count / elapsed:.0f} images/sec\")\n    print(f\"  Cache: {CACHE_IMAGES_PATH} \"\n          f\"({CACHE_IMAGES_PATH.stat().st_size / 1e9:.2f} GB)\")\n\n    # Sanity check: all rows were written (no image_id in train.csv missed)\n    assert processed_count == N_TOTAL_EXPECTED, \\\n        f\"Processed {processed_count} but expected {N_TOTAL_EXPECTED} — some image_ids missing?\"\n\n    # Save labels in the same row order as the mmap\n    labels_df = df[['image_id', 'grapheme_root', 'vowel_diacritic',\n                    'consonant_diacritic', 'grapheme']].copy()\n    labels_df.to_parquet(CACHE_LABELS_PATH, index=False)\n    print(f\"✓ Labels saved to {CACHE_LABELS_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:08:08.614248Z","iopub.execute_input":"2026-09-19T06:08:08.614966Z","iopub.status.idle":"2026-09-19T06:10:14.189237Z","shell.execute_reply.started":"2026-09-19T06:08:08.614937Z","shell.execute_reply":"2026-09-19T06:10:14.188606Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 5: Verify cache — memmap load + random-sample grid\nimages = np.load(CACHE_IMAGES_PATH, mmap_mode='r')\nlabels_df = pd.read_parquet(CACHE_LABELS_PATH)\n\nprint(f\"✓ Loaded cache\")\nprint(f\"  Images shape: {images.shape}, dtype: {images.dtype}\")\nprint(f\"  Labels:       {len(labels_df):,} rows\")\n\n# Structural assertions\nassert images.shape == (N_TOTAL_EXPECTED, TARGET_SIZE, TARGET_SIZE)\nassert len(labels_df) == N_TOTAL_EXPECTED\nassert images.dtype == np.uint8\n\n# Row alignment: labels_df.iloc[i] should describe images[i]\nassert (labels_df['image_id'].values == df['image_id'].values).all(), \\\n    \"labels_df is not in the same order as train.csv — row alignment broken!\"\n\n# No blank images\nfirst_thousand_max = images[:1000].reshape(1000, -1).max(axis=1)\nn_blank = int((first_thousand_max == 0).sum())\nassert n_blank == 0, f\"{n_blank}/1000 processed images appear fully blank\"\n\nprint(f\"  ✓ Shape / dtype / order / non-blank checks pass\")\n\n# 16 random samples with labels\nrng = np.random.default_rng(42)\nsample_indices = rng.choice(N_TOTAL_EXPECTED, size=16, replace=False)\n\nfig, axes = plt.subplots(4, 4, figsize=(12, 13))\nfor ax, idx in zip(axes.flat, sample_indices):\n    img = images[idx]\n    row = labels_df.iloc[idx]\n    ax.imshow(img, cmap='gray')\n    ax.set_title(f\"{row['grapheme']}\\n\"\n                 f\"root={row['grapheme_root']}  \"\n                 f\"v={row['vowel_diacritic']}  \"\n                 f\"c={row['consonant_diacritic']}\",\n                 fontsize=10)\n    ax.axis('off')\nplt.suptitle('16 random cached samples (128×128, bright ink on dark)',\n             fontsize=13, y=1.00)\nplt.tight_layout()\nplt.savefig(FIG_DIR / 'eda_samples.png', dpi=140, bbox_inches='tight')\nplt.show()\n\n# Dataset summary for the slides\nprint(f\"\\n--- Dataset summary (for slides) ---\")\nprint(f\"  Total images:           {N_TOTAL_EXPECTED:,}\")\nprint(f\"  Image size (cached):    {TARGET_SIZE} × {TARGET_SIZE} grayscale uint8\")\nprint(f\"  Cache footprint:        {CACHE_IMAGES_PATH.stat().st_size / 1e9:.2f} GB\")\nprint(f\"  Unique grapheme wholes: {labels_df['grapheme'].nunique():,}\")\nprint(f\"  Class counts:           {N_ROOT_CLASSES} roots / \"\n      f\"{N_VOWEL_CLASSES} vowels / {N_CONSONANT_CLASSES} consonants\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:10:50.433545Z","iopub.execute_input":"2026-09-19T06:10:50.43399Z","iopub.status.idle":"2026-09-19T06:10:52.406363Z","shell.execute_reply.started":"2026-09-19T06:10:50.433961Z","shell.execute_reply":"2026-09-19T06:10:52.405652Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Phase 2**","metadata":{}},{"cell_type":"code","source":"# Phase 2, Cell 1: Imports and reload cache from Phase 1\nimport os\nimport gc\nimport json\nimport time\nfrom pathlib import Path\nfrom collections import Counter\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# Paths (same as Phase 1)\nWORK_DIR = Path('/kaggle/working')\nCACHE_DIR = WORK_DIR / 'cache'\nSPLIT_DIR = WORK_DIR / 'splits'\nFIG_DIR = WORK_DIR / 'figures'\nSPLIT_DIR.mkdir(parents=True, exist_ok=True)\n\nCACHE_IMAGES_PATH = CACHE_DIR / 'images_128.npy'\nCACHE_LABELS_PATH = CACHE_DIR / 'labels.parquet'\n\n# Constants\nTARGET_SIZE = 128\nN_ROOT = 168\nN_VOWEL = 11\nN_CONS = 7\n\n# Load cache\nimages = np.load(CACHE_IMAGES_PATH, mmap_mode='r')\nlabels_df = pd.read_parquet(CACHE_LABELS_PATH)\n\nprint(f\"✓ Cache loaded\")\nprint(f\"  Images: {images.shape}, dtype={images.dtype}\")\nprint(f\"  Labels: {len(labels_df):,} rows\")\nprint(f\"  Unique wholes: {labels_df['grapheme'].nunique():,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:22:23.300775Z","iopub.execute_input":"2026-09-19T06:22:23.301507Z","iopub.status.idle":"2026-09-19T06:22:30.469618Z","shell.execute_reply.started":"2026-09-19T06:22:23.301477Z","shell.execute_reply":"2026-09-19T06:22:30.46897Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 2, Cell 2 (FIXED): Split-by-whole with component-coverage guarantee\n\ndef build_split_by_whole(labels_df, holdout_frac, seed=42):\n    \"\"\"Split the dataset by grapheme WHOLE, not by image.\n    \n    Guarantees that after holding out wholes, every root, vowel, and\n    consonant class still has at least MIN_KEEP wholes remaining in \n    the seen pool (train+val+seen_test).\n    \n    Returns split dict + metadata.\n    \"\"\"\n    MIN_KEEP = 2  # each component must keep at least 2 wholes in seen pool\n    \n    rng = np.random.default_rng(seed)\n    \n    # 1. List all unique wholes\n    wholes = labels_df['grapheme'].unique().tolist()\n    rng.shuffle(wholes)\n    \n    # 2. For each whole, record its components\n    whole_info = {}\n    for g in wholes:\n        row = labels_df[labels_df['grapheme'] == g].iloc[0]\n        whole_info[g] = {\n            'root': int(row['grapheme_root']),\n            'vowel': int(row['vowel_diacritic']),\n            'cons': int(row['consonant_diacritic']),\n        }\n    \n    # 3. Count how many wholes each component value appears in\n    from collections import Counter\n    root_counts = Counter(info['root'] for info in whole_info.values())\n    vowel_counts = Counter(info['vowel'] for info in whole_info.values())\n    cons_counts = Counter(info['cons'] for info in whole_info.values())\n    \n    # Print rare components for awareness\n    print(\"  Rare components (appear in ≤3 wholes):\")\n    for name, counts in [('root', root_counts), ('vowel', vowel_counts), ('cons', cons_counts)]:\n        rare = {k: v for k, v in counts.items() if v <= 3}\n        if rare:\n            print(f\"    {name}: {dict(sorted(rare.items()))}\")\n    \n    # 4. Greedily select wholes for hold-out, skipping any that would\n    #    drop a component below MIN_KEEP wholes in the seen pool\n    n_target = int(holdout_frac * len(wholes))\n    \n    # Working counters: start with full counts, decrement as we hold out\n    remaining_root = dict(root_counts)\n    remaining_vowel = dict(vowel_counts)\n    remaining_cons = dict(cons_counts)\n    \n    unseen_wholes = []\n    \n    for g in wholes:\n        if len(unseen_wholes) >= n_target:\n            break\n        \n        info = whole_info[g]\n        r, v, c = info['root'], info['vowel'], info['cons']\n        \n        # Check: would holding this whole drop any component below MIN_KEEP?\n        if (remaining_root[r] <= MIN_KEEP or\n            remaining_vowel[v] <= MIN_KEEP or\n            remaining_cons[c] <= MIN_KEEP):\n            # Skip — this whole is needed to keep coverage\n            continue\n        \n        # Safe to hold out\n        unseen_wholes.append(g)\n        remaining_root[r] -= 1\n        remaining_vowel[v] -= 1\n        remaining_cons[c] -= 1\n    \n    unseen_wholes_set = set(unseen_wholes)\n    actual_holdout = len(unseen_wholes) / len(wholes)\n    \n    print(f\"  Requested hold-out: {holdout_frac:.0%} ({n_target} wholes)\")\n    print(f\"  Actual hold-out:    {actual_holdout:.1%} ({len(unseen_wholes)} wholes)\")\n    if len(unseen_wholes) < n_target:\n        print(f\"  ({n_target - len(unseen_wholes)} wholes skipped to preserve coverage)\")\n    \n    # 5. All images of unseen wholes → unseen_test\n    is_unseen = labels_df['grapheme'].isin(unseen_wholes_set)\n    unseen_indices = labels_df.index[is_unseen].tolist()\n    \n    # 6. Remaining → stratified split into train / val / seen_test\n    pool_df = labels_df[~is_unseen].copy()\n    pool_indices = pool_df.index.values\n    strat_col = pool_df['grapheme_root'].values\n    \n    train_idx, temp_idx = train_test_split(\n        pool_indices, test_size=0.20, stratify=strat_col, random_state=seed\n    )\n    temp_strat = labels_df.loc[temp_idx, 'grapheme_root'].values\n    val_idx, seen_test_idx = train_test_split(\n        temp_idx, test_size=0.50, stratify=temp_strat, random_state=seed\n    )\n    \n    split = {\n        'train':       train_idx.tolist(),\n        'val':         val_idx.tolist(),\n        'seen_test':   seen_test_idx.tolist(),\n        'unseen_test': unseen_indices,\n    }\n    \n    meta = {\n        'seed': seed,\n        'holdout_frac_requested': holdout_frac,\n        'holdout_frac_actual': round(actual_holdout, 4),\n        'n_wholes_total': len(wholes),\n        'n_wholes_unseen': len(unseen_wholes),\n        'n_wholes_seen': len(wholes) - len(unseen_wholes),\n        'unseen_wholes': sorted(unseen_wholes),\n        'split_sizes': {k: len(v) for k, v in split.items()},\n    }\n    \n    return split, meta\n\n\n# Build the main split (20% hold-out)\nprint(\"=== Building main split (20% hold-out) ===\")\nsplit_main, meta_main = build_split_by_whole(labels_df, holdout_frac=0.20, seed=42)\n\nprint(f\"\\n=== Main split summary ===\")\nfor k, v in meta_main['split_sizes'].items():\n    print(f\"  {k:15s}: {v:>7,} images\")\nprint(f\"  Total:          {sum(meta_main['split_sizes'].values()):>7,}\")\nprint(f\"  Wholes total:   {meta_main['n_wholes_total']}\")\nprint(f\"  Wholes unseen:  {meta_main['n_wholes_unseen']}\")\nprint(f\"  Wholes seen:    {meta_main['n_wholes_seen']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:24:20.712029Z","iopub.execute_input":"2026-09-19T06:24:20.712578Z","iopub.status.idle":"2026-09-19T06:24:38.884955Z","shell.execute_reply.started":"2026-09-19T06:24:20.712536Z","shell.execute_reply":"2026-09-19T06:24:38.884202Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 2, Cell 3: Leakage guards and coverage checks\n\ndef verify_split(split, labels_df, meta):\n    \"\"\"Run every check that could catch a broken split.\"\"\"\n    \n    train_set = set(split['train'])\n    val_set = set(split['val'])\n    seen_set = set(split['seen_test'])\n    unseen_set = set(split['unseen_test'])\n    all_sets = [train_set, val_set, seen_set, unseen_set]\n    names = ['train', 'val', 'seen_test', 'unseen_test']\n    \n    # 1. No overlap between any pair\n    for i in range(len(all_sets)):\n        for j in range(i + 1, len(all_sets)):\n            overlap = all_sets[i] & all_sets[j]\n            assert len(overlap) == 0, \\\n                f\"LEAK: {names[i]} ∩ {names[j]} has {len(overlap)} shared indices\"\n    print(\"✓ No index overlap between any split pair\")\n    \n    # 2. All indices accounted for\n    total = sum(len(s) for s in all_sets)\n    assert total == len(labels_df), \\\n        f\"Split has {total} indices but labels_df has {len(labels_df)}\"\n    print(f\"✓ All {total:,} indices accounted for\")\n    \n    # 3. No unseen WHOLE leaks into train\n    unseen_wholes = set(meta['unseen_wholes'])\n    train_wholes = set(labels_df.iloc[split['train']]['grapheme'].unique())\n    leaked = train_wholes & unseen_wholes\n    assert len(leaked) == 0, \\\n        f\"LEAK: {len(leaked)} unseen wholes found in train: {list(leaked)[:5]}\"\n    print(f\"✓ No unseen whole appears in train\")\n    \n    # 4. Every COMPONENT still appears in train\n    train_roots = set(labels_df.iloc[split['train']]['grapheme_root'].unique())\n    train_vowels = set(labels_df.iloc[split['train']]['vowel_diacritic'].unique())\n    train_cons = set(labels_df.iloc[split['train']]['consonant_diacritic'].unique())\n    \n    assert train_roots == set(range(N_ROOT)), \\\n        f\"Missing roots in train: {set(range(N_ROOT)) - train_roots}\"\n    assert train_vowels == set(range(N_VOWEL)), \\\n        f\"Missing vowels in train: {set(range(N_VOWEL)) - train_vowels}\"\n    assert train_cons == set(range(N_CONS)), \\\n        f\"Missing consonants in train: {set(range(N_CONS)) - train_cons}\"\n    print(f\"✓ All {N_ROOT} roots, {N_VOWEL} vowels, {N_CONS} consonants present in train\")\n    \n    # 5. Unseen test contains ONLY unseen wholes\n    unseen_test_wholes = set(labels_df.iloc[split['unseen_test']]['grapheme'].unique())\n    non_unseen_in_test = unseen_test_wholes - unseen_wholes\n    assert len(non_unseen_in_test) == 0, \\\n        f\"Unseen test contains {len(non_unseen_in_test)} wholes NOT in unseen set\"\n    print(f\"✓ Unseen test contains only unseen wholes ({len(unseen_test_wholes)} wholes)\")\n    \n    # 6. Val and seen_test contain only SEEN wholes\n    val_wholes = set(labels_df.iloc[split['val']]['grapheme'].unique())\n    seen_test_wholes = set(labels_df.iloc[split['seen_test']]['grapheme'].unique())\n    assert val_wholes.issubset(train_wholes | val_wholes), \"Val has wholes outside seen set\"\n    assert len(val_wholes & unseen_wholes) == 0, \"Val contains unseen wholes!\"\n    assert len(seen_test_wholes & unseen_wholes) == 0, \"Seen test contains unseen wholes!\"\n    print(f\"✓ Val and seen_test contain only seen wholes\")\n    \n    print(\"\\n✓ ALL LEAKAGE GUARDS PASS\")\n\n\nverify_split(split_main, labels_df, meta_main)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:24:51.237804Z","iopub.execute_input":"2026-09-19T06:24:51.238604Z","iopub.status.idle":"2026-09-19T06:24:51.353815Z","shell.execute_reply.started":"2026-09-19T06:24:51.238577Z","shell.execute_reply":"2026-09-19T06:24:51.352958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 2, Cell 4: Build dose-response splits (10/20/30/40%) and save all\n\nholdout_fracs = [0.10, 0.20, 0.30, 0.40]\n\nfor frac in holdout_fracs:\n    split, meta = build_split_by_whole(labels_df, holdout_frac=frac, seed=42)\n    \n    # Verify every split\n    print(f\"\\n--- Holdout {frac:.0%} ---\")\n    verify_split(split, labels_df, meta)\n    \n    # Save split indices + metadata\n    save_path = SPLIT_DIR / f'split_holdout_{int(frac*100)}pct_seed42.json'\n    save_data = {\n        'meta': meta,\n        'indices': {k: v for k, v in split.items()},\n    }\n    # Convert unseen_wholes list to saveable format (already sorted strings)\n    with open(save_path, 'w') as f:\n        json.dump(save_data, f, indent=2, ensure_ascii=False)\n    print(f\"  ✓ Saved to {save_path.name}\")\n\n# List saved splits\nprint(f\"\\n=== Saved split files ===\")\nfor p in sorted(SPLIT_DIR.glob('*.json')):\n    print(f\"  {p.name}  ({p.stat().st_size / 1e3:.1f} KB)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:25:39.283394Z","iopub.execute_input":"2026-09-19T06:25:39.283869Z","iopub.status.idle":"2026-09-19T06:26:52.295663Z","shell.execute_reply.started":"2026-09-19T06:25:39.283838Z","shell.execute_reply":"2026-09-19T06:26:52.295068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 2, Cell 5: Visualize the main split\n\n# Reload the main split\nwith open(SPLIT_DIR / 'split_holdout_20pct_seed42.json', 'r') as f:\n    saved = json.load(f)\nsplit_main = saved['indices']\nmeta_main = saved['meta']\n\n# Bar chart: split sizes\nsplits_names = ['train', 'val', 'seen_test', 'unseen_test']\nsplits_sizes = [len(split_main[k]) for k in splits_names]\ncolors = ['#16203A', '#5A6478', '#1F6F6B', '#C0392B']\n\nfig, axes = plt.subplots(1, 3, figsize=(18, 4.5))\n\n# Panel 1: Split sizes\nax = axes[0]\nbars = ax.bar(splits_names, splits_sizes, color=colors)\nfor bar, size in zip(bars, splits_sizes):\n    ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 500,\n            f'{size:,}', ha='center', fontsize=11)\nax.set_title('Split sizes (20% hold-out)', fontsize=13)\nax.set_ylabel('Number of images')\nax.grid(axis='y', alpha=0.3)\n\n# Panel 2: Number of unique wholes per split\nwholes_per_split = []\nfor k in splits_names:\n    w = labels_df.iloc[split_main[k]]['grapheme'].nunique()\n    wholes_per_split.append(w)\nax = axes[1]\nbars = ax.bar(splits_names, wholes_per_split, color=colors)\nfor bar, w in zip(bars, wholes_per_split):\n    ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 3,\n            str(w), ha='center', fontsize=11)\nax.set_title('Unique wholes per split', fontsize=13)\nax.set_ylabel('Number of unique grapheme wholes')\nax.grid(axis='y', alpha=0.3)\n\n# Panel 3: Root class distribution across train vs unseen\nax = axes[2]\ntrain_root_dist = labels_df.iloc[split_main['train']]['grapheme_root'].value_counts().sort_index()\nunseen_root_dist = labels_df.iloc[split_main['unseen_test']]['grapheme_root'].value_counts().sort_index()\n# Align indices (some roots might have 0 in unseen)\nall_roots = range(N_ROOT)\ntrain_vals = [train_root_dist.get(r, 0) for r in all_roots]\nunseen_vals = [unseen_root_dist.get(r, 0) for r in all_roots]\nx = np.arange(N_ROOT)\nax.bar(x - 0.2, train_vals, width=0.4, color='#16203A', label='Train', alpha=0.8)\nax.bar(x + 0.2, unseen_vals, width=0.4, color='#C0392B', label='Unseen test', alpha=0.8)\nax.set_title('Root distribution: train vs unseen', fontsize=13)\nax.set_xlabel('Root class')\nax.set_ylabel('Count')\nax.set_yscale('log')\nax.legend()\nax.grid(axis='y', alpha=0.3)\n\nplt.suptitle('Split-by-whole: structure of the 20% hold-out split', fontsize=14, y=1.02)\nplt.tight_layout()\nplt.savefig(FIG_DIR / 'split_structure.png', dpi=140, bbox_inches='tight')\nplt.show()\n\n# Key observation for the slides\ntrain_wholes = set(labels_df.iloc[split_main['train']]['grapheme'].unique())\nunseen_wholes = set(labels_df.iloc[split_main['unseen_test']]['grapheme'].unique())\nprint(f\"\\n--- Key numbers for slides ---\")\nprint(f\"  Train wholes:        {len(train_wholes)}\")\nprint(f\"  Unseen wholes:       {len(unseen_wholes)}\")\nprint(f\"  Overlap (must be 0): {len(train_wholes & unseen_wholes)}\")\nprint(f\"  Components shared:   all {N_ROOT} roots, {N_VOWEL} vowels, {N_CONS} consonants\")\nprint(f\"  → The parts are familiar; only the combination is new.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:28:02.906108Z","iopub.execute_input":"2026-09-19T06:28:02.906717Z","iopub.status.idle":"2026-09-19T06:28:04.939521Z","shell.execute_reply.started":"2026-09-19T06:28:02.906643Z","shell.execute_reply":"2026-09-19T06:28:04.938888Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 2, Cell 6: PyTorch Dataset\n\nclass BengaliDataset(Dataset):\n    \"\"\"Dataset for Bengali grapheme recognition.\n    \n    Reads from the memmap cache (microsecond loads).\n    Returns both component labels (root/vowel/cons) and a flat \n    whole-grapheme index (for the flat-classifier baseline).\n    \"\"\"\n    \n    def __init__(self, images_mmap, labels_df, indices, transform=None,\n                 whole_to_idx=None):\n        \"\"\"\n        Args:\n            images_mmap:  np.memmap of shape (N, 128, 128) uint8\n            labels_df:    DataFrame with columns [image_id, grapheme_root,\n                          vowel_diacritic, consonant_diacritic, grapheme]\n            indices:      list of integer indices into images_mmap / labels_df\n            transform:    albumentations transform (or None for val/test)\n            whole_to_idx: dict mapping grapheme string → flat class index.\n                          Built from training set only. Unseen wholes get idx -1.\n        \"\"\"\n        self.images = images_mmap\n        self.indices = np.array(indices, dtype=np.int64)\n        \n        # Pre-extract labels for speed (avoid DataFrame lookup in __getitem__)\n        sub = labels_df.iloc[indices]\n        self.roots = sub['grapheme_root'].values.astype(np.int64)\n        self.vowels = sub['vowel_diacritic'].values.astype(np.int64)\n        self.cons = sub['consonant_diacritic'].values.astype(np.int64)\n        self.graphemes = sub['grapheme'].values\n        \n        self.transform = transform\n        self.whole_to_idx = whole_to_idx or {}\n    \n    def __len__(self):\n        return len(self.indices)\n    \n    def __getitem__(self, i):\n        # Load image from memmap (microsecond)\n        idx = self.indices[i]\n        img = self.images[idx].copy()  # copy from read-only memmap\n        \n        # Albumentations expects (H, W, C)\n        img = img[:, :, np.newaxis]  # (128, 128, 1)\n        \n        if self.transform:\n            augmented = self.transform(image=img)\n            img = augmented['image']  # returns (C, H, W) tensor if ToTensorV2 is in the pipeline\n        else:\n            # Manual: normalize and convert\n            img = img.astype(np.float32) / 255.0\n            img = torch.from_numpy(img).permute(2, 0, 1)  # (1, 128, 128)\n        \n        # Labels\n        root = torch.tensor(self.roots[i], dtype=torch.long)\n        vowel = torch.tensor(self.vowels[i], dtype=torch.long)\n        cons = torch.tensor(self.cons[i], dtype=torch.long)\n        \n        # Flat whole-grapheme index (-1 if unseen)\n        grapheme_str = self.graphemes[i]\n        whole_idx = self.whole_to_idx.get(grapheme_str, -1)\n        whole_idx = torch.tensor(whole_idx, dtype=torch.long)\n        \n        return img, {\n            'root': root,\n            'vowel': vowel,\n            'cons': cons,\n            'whole_idx': whole_idx,\n        }\n\n\ndef build_whole_to_idx(labels_df, train_indices):\n    \"\"\"Build the flat-classifier's class mapping from the training set only.\n    \n    Returns:\n        whole_to_idx: dict mapping grapheme string → integer index (0..N-1)\n        idx_to_components: dict mapping integer index → (root, vowel, cons)\n            Used later to decompose flat predictions into components for scoring.\n    \"\"\"\n    train_graphemes = labels_df.iloc[train_indices]['grapheme'].unique()\n    train_graphemes = sorted(train_graphemes)  # deterministic ordering\n    \n    whole_to_idx = {g: i for i, g in enumerate(train_graphemes)}\n    \n    # For each whole, store its component labels (for scoring the flat model)\n    idx_to_components = {}\n    for g in train_graphemes:\n        row = labels_df[labels_df['grapheme'] == g].iloc[0]\n        idx_to_components[whole_to_idx[g]] = {\n            'root': int(row['grapheme_root']),\n            'vowel': int(row['vowel_diacritic']),\n            'cons': int(row['consonant_diacritic']),\n        }\n    \n    return whole_to_idx, idx_to_components\n\n\n# Build the mapping from the main split's training set\nwhole_to_idx, idx_to_components = build_whole_to_idx(labels_df, split_main['train'])\n\nprint(f\"✓ BengaliDataset class defined\")\nprint(f\"✓ Flat-classifier mapping built\")\nprint(f\"  Seen wholes (flat output classes): {len(whole_to_idx)}\")\nprint(f\"  idx_to_components entries:         {len(idx_to_components)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:28:19.138115Z","iopub.execute_input":"2026-09-19T06:28:19.138515Z","iopub.status.idle":"2026-09-19T06:28:33.611072Z","shell.execute_reply.started":"2026-09-19T06:28:19.138486Z","shell.execute_reply":"2026-09-19T06:28:33.610359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 2, Cell 7: Augmentation pipelines\n\ndef get_train_transform(image_size=TARGET_SIZE):\n    \"\"\"Training augmentation: geometric + dropout. \n    No horizontal flips — flipping is invalid for characters.\n    \"\"\"\n    return A.Compose([\n        A.ShiftScaleRotate(\n            shift_limit=0.06, scale_limit=0.10, rotate_limit=15,\n            border_mode=0, value=0, p=0.7\n        ),\n        A.OneOf([\n            A.ElasticTransform(alpha=30, sigma=5, p=1.0),\n            A.GridDistortion(num_steps=5, distort_limit=0.1, p=1.0),\n            A.OpticalDistortion(distort_limit=0.1, shift_limit=0.1, p=1.0),\n        ], p=0.3),\n        A.CoarseDropout(\n            max_holes=4, max_height=16, max_width=16,\n            min_holes=1, min_height=8, min_width=8,\n            fill_value=0, p=0.3\n        ),\n        # Normalize: we use single-channel; for pretrained backbones we'll\n        # replicate to 3 channels in the model, not here\n        A.Normalize(mean=[0.5], std=[0.5]),  # maps [0,1] → [-1,1]\n        ToTensorV2(),\n    ])\n\n\ndef get_val_transform(image_size=TARGET_SIZE):\n    \"\"\"Validation/test: normalize only.\"\"\"\n    return A.Compose([\n        A.Normalize(mean=[0.5], std=[0.5]),\n        ToTensorV2(),\n    ])\n\n\n# Quick visual test: apply train augmentation 8 times to the same image\ntest_idx = split_main['train'][0]\ntest_img = images[test_idx]\ntrain_tf = get_train_transform()\n\nfig, axes = plt.subplots(2, 5, figsize=(16, 6.5))\n# First column: original (val transform)\nval_tf = get_val_transform()\norig = val_tf(image=test_img[:, :, np.newaxis])['image']\naxes[0, 0].imshow(orig.squeeze(), cmap='gray')\naxes[0, 0].set_title('Original', fontsize=11)\naxes[0, 0].axis('off')\naxes[1, 0].imshow(orig.squeeze(), cmap='gray')\naxes[1, 0].set_title('Original', fontsize=11)\naxes[1, 0].axis('off')\n\n# Remaining: 8 random augmentations\nfor ax_idx, ax in enumerate(list(axes[0, 1:]) + list(axes[1, 1:])):\n    aug = train_tf(image=test_img[:, :, np.newaxis])['image']\n    ax.imshow(aug.squeeze(), cmap='gray')\n    ax.set_title(f'Aug #{ax_idx + 1}', fontsize=10)\n    ax.axis('off')\n\nplt.suptitle('Train augmentation: shift/scale/rotate, elastic/grid/optical distortion, coarse dropout',\n             fontsize=12, y=1.01)\nplt.tight_layout()\nplt.savefig(FIG_DIR / 'augmentation_samples.png', dpi=140, bbox_inches='tight')\nplt.show()\n\nprint(\"✓ Augmentation pipelines defined and tested\")\nprint(\"  NOTE: No horizontal flip — flipping is invalid for characters.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:28:50.4422Z","iopub.execute_input":"2026-09-19T06:28:50.442512Z","iopub.status.idle":"2026-09-19T06:28:51.653889Z","shell.execute_reply.started":"2026-09-19T06:28:50.442484Z","shell.execute_reply":"2026-09-19T06:28:51.65306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 2, Cell 8: CutMix and MixUp (applied at batch level during training)\n\ndef cutmix_batch(images, labels, alpha=1.0):\n    \"\"\"Apply CutMix to a batch.\n    \n    Args:\n        images: (B, C, H, W) tensor\n        labels: dict with 'root', 'vowel', 'cons' each (B,) long tensor\n    \n    Returns:\n        mixed_images, labels_a, labels_b, lam (float)\n    \"\"\"\n    B = images.size(0)\n    lam = np.random.beta(alpha, alpha)\n    \n    # Random permutation for pairing\n    perm = torch.randperm(B)\n    \n    # Random bounding box\n    _, _, H, W = images.shape\n    cut_ratio = np.sqrt(1.0 - lam)\n    cut_h = int(H * cut_ratio)\n    cut_w = int(W * cut_ratio)\n    cy = np.random.randint(H)\n    cx = np.random.randint(W)\n    y1 = max(0, cy - cut_h // 2)\n    y2 = min(H, cy + cut_h // 2)\n    x1 = max(0, cx - cut_w // 2)\n    x2 = min(W, cx + cut_w // 2)\n    \n    # Mix\n    mixed = images.clone()\n    mixed[:, :, y1:y2, x1:x2] = images[perm, :, y1:y2, x1:x2]\n    \n    # Adjust lambda to actual area ratio\n    lam = 1.0 - (y2 - y1) * (x2 - x1) / (H * W)\n    \n    labels_a = labels\n    labels_b = {k: v[perm] for k, v in labels.items()}\n    \n    return mixed, labels_a, labels_b, lam\n\n\ndef mixup_batch(images, labels, alpha=0.4):\n    \"\"\"Apply MixUp to a batch.\n    \n    Args:\n        images: (B, C, H, W) tensor\n        labels: dict with 'root', 'vowel', 'cons' each (B,) long tensor\n    \n    Returns:\n        mixed_images, labels_a, labels_b, lam (float)\n    \"\"\"\n    B = images.size(0)\n    lam = np.random.beta(alpha, alpha)\n    lam = max(lam, 1.0 - lam)  # ensure lam >= 0.5 so label_a dominates\n    \n    perm = torch.randperm(B)\n    \n    mixed = lam * images + (1.0 - lam) * images[perm]\n    \n    labels_a = labels\n    labels_b = {k: v[perm] for k, v in labels.items()}\n    \n    return mixed, labels_a, labels_b, lam\n\n\ndef mix_criterion(criterion, pred, labels_a, labels_b, lam):\n    \"\"\"Compute mixed loss for CutMix/MixUp.\n    \n    criterion: function(pred_dict, label_dict) → scalar loss\n    pred: dict with 'root', 'vowel', 'cons' logits\n    labels_a, labels_b: dicts with 'root', 'vowel', 'cons' targets\n    lam: mixing ratio\n    \n    Returns: lam * loss(pred, labels_a) + (1 - lam) * loss(pred, labels_b)\n    \"\"\"\n    return lam * criterion(pred, labels_a) + (1.0 - lam) * criterion(pred, labels_b)\n\n\n# Quick test\ntest_imgs = torch.randn(8, 1, 128, 128)\ntest_labels = {\n    'root': torch.randint(0, N_ROOT, (8,)),\n    'vowel': torch.randint(0, N_VOWEL, (8,)),\n    'cons': torch.randint(0, N_CONS, (8,)),\n}\n\nmixed, la, lb, lam = cutmix_batch(test_imgs, test_labels, alpha=1.0)\nprint(f\"✓ CutMix test:  output shape={mixed.shape}, lam={lam:.3f}\")\n\nmixed, la, lb, lam = mixup_batch(test_imgs, test_labels, alpha=0.4)\nprint(f\"✓ MixUp test:   output shape={mixed.shape}, lam={lam:.3f}\")\nprint(f\"✓ mix_criterion ready (used in training loop)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:29:05.147749Z","iopub.execute_input":"2026-09-19T06:29:05.148171Z","iopub.status.idle":"2026-09-19T06:29:05.217091Z","shell.execute_reply.started":"2026-09-19T06:29:05.148142Z","shell.execute_reply":"2026-09-19T06:29:05.216407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 2, Cell 9: DataLoader smoke test — pull one real batch\n\n# Build datasets for the main split\ntrain_ds = BengaliDataset(\n    images_mmap=images,\n    labels_df=labels_df,\n    indices=split_main['train'],\n    transform=get_train_transform(),\n    whole_to_idx=whole_to_idx,\n)\n\nval_ds = BengaliDataset(\n    images_mmap=images,\n    labels_df=labels_df,\n    indices=split_main['val'],\n    transform=get_val_transform(),\n    whole_to_idx=whole_to_idx,\n)\n\nunseen_ds = BengaliDataset(\n    images_mmap=images,\n    labels_df=labels_df,\n    indices=split_main['unseen_test'],\n    transform=get_val_transform(),\n    whole_to_idx=whole_to_idx,\n)\n\nprint(f\"Datasets created:\")\nprint(f\"  Train:       {len(train_ds):>7,} images\")\nprint(f\"  Val:         {len(val_ds):>7,} images\")\nprint(f\"  Unseen test: {len(unseen_ds):>7,} images\")\n\n# Build DataLoaders\ntrain_dl = DataLoader(train_ds, batch_size=32, shuffle=True,\n                      num_workers=2, pin_memory=True, drop_last=True)\nval_dl = DataLoader(val_ds, batch_size=64, shuffle=False,\n                    num_workers=2, pin_memory=True)\nunseen_dl = DataLoader(unseen_ds, batch_size=64, shuffle=False,\n                       num_workers=2, pin_memory=True)\n\n# Pull one batch from each\nfor name, dl in [('Train', train_dl), ('Val', val_dl), ('Unseen', unseen_dl)]:\n    batch_img, batch_labels = next(iter(dl))\n    print(f\"\\n--- {name} batch ---\")\n    print(f\"  Image:     shape={batch_img.shape}, dtype={batch_img.dtype}, \"\n          f\"range=[{batch_img.min():.2f}, {batch_img.max():.2f}]\")\n    print(f\"  Root:      shape={batch_labels['root'].shape}, \"\n          f\"range=[{batch_labels['root'].min()}, {batch_labels['root'].max()}]\")\n    print(f\"  Vowel:     shape={batch_labels['vowel'].shape}, \"\n          f\"range=[{batch_labels['vowel'].min()}, {batch_labels['vowel'].max()}]\")\n    print(f\"  Cons:      shape={batch_labels['cons'].shape}, \"\n          f\"range=[{batch_labels['cons'].min()}, {batch_labels['cons'].max()}]\")\n    print(f\"  Whole idx: shape={batch_labels['whole_idx'].shape}, \"\n          f\"range=[{batch_labels['whole_idx'].min()}, {batch_labels['whole_idx'].max()}]\")\n    \n    # Verify label ranges\n    assert batch_labels['root'].min() >= 0 and batch_labels['root'].max() < N_ROOT\n    assert batch_labels['vowel'].min() >= 0 and batch_labels['vowel'].max() < N_VOWEL\n    assert batch_labels['cons'].min() >= 0 and batch_labels['cons'].max() < N_CONS\n\n# Special check: unseen test should have whole_idx == -1 for many samples\nunseen_batch_img, unseen_batch_labels = next(iter(unseen_dl))\nn_unknown = (unseen_batch_labels['whole_idx'] == -1).sum().item()\nprint(f\"\\n--- Unseen flat-idx check ---\")\nprint(f\"  Samples with whole_idx == -1: {n_unknown}/{len(unseen_batch_labels['whole_idx'])}\")\nprint(f\"  (Expected: ALL of them, since these wholes never appeared in train)\")\n\n# Visualize one train batch (augmented)\ntrain_batch_img, train_batch_labels = next(iter(train_dl))\nfig, axes = plt.subplots(2, 8, figsize=(18, 5))\nfor i, ax in enumerate(axes.flat):\n    if i >= train_batch_img.size(0):\n        ax.axis('off')\n        continue\n    img = train_batch_img[i].squeeze().numpy()\n    # Undo normalization for display: x * std + mean = x * 0.5 + 0.5\n    img = img * 0.5 + 0.5\n    ax.imshow(img, cmap='gray', vmin=0, vmax=1)\n    r = train_batch_labels['root'][i].item()\n    v = train_batch_labels['vowel'][i].item()\n    c = train_batch_labels['cons'][i].item()\n    ax.set_title(f'r={r} v={v} c={c}', fontsize=9)\n    ax.axis('off')\nplt.suptitle('One training batch (augmented, normalized)', fontsize=12, y=1.01)\nplt.tight_layout()\nplt.show()\n\nprint(\"\\n✓ Phase 2 complete — splits verified, dataset and dataloaders working\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T06:30:36.398487Z","iopub.execute_input":"2026-09-19T06:30:36.399084Z","iopub.status.idle":"2026-09-19T06:30:38.896595Z","shell.execute_reply.started":"2026-09-19T06:30:36.399053Z","shell.execute_reply":"2026-09-19T06:30:38.895533Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Phase 3 — Model, Loss, and Metrics","metadata":{}},{"cell_type":"code","source":"# Phase 3, Cell 1 (fixed): Imports\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport numpy as np\nfrom sklearn.metrics import recall_score\nfrom pathlib import Path\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Device: {device}\")\nif device.type == 'cuda':\n    print(f\"  GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"  Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB\")\n\n# Constants (carried from Phase 2)\nN_ROOT = 168\nN_VOWEL = 11\nN_CONS = 7\nTARGET_SIZE = 128","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:17:30.469206Z","iopub.execute_input":"2026-09-19T07:17:30.469831Z","iopub.status.idle":"2026-09-19T07:17:30.476849Z","shell.execute_reply.started":"2026-09-19T07:17:30.4698Z","shell.execute_reply":"2026-09-19T07:17:30.47598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 3, Cell 2: Generalized Mean (GeM) pooling\n\nclass GeM(nn.Module):\n    \"\"\"Generalized Mean pooling.\n    \n    At p=1 → Global Average Pooling (GAP)\n    At p→∞ → Global Max Pooling (GMP)\n    p is learnable — the model decides the best trade-off.\n    \n    Why we use it: handwriting has strong local features (strokes, curves).\n    GAP dilutes them by averaging with background. GMP keeps only the peak\n    and loses distributed info. GeM's learned p concentrates weight on\n    strong activations while retaining weaker ones.\n    \"\"\"\n    def __init__(self, p=3.0, eps=1e-6):\n        super().__init__()\n        self.p = nn.Parameter(torch.tensor(float(p)))\n        self.eps = eps\n    \n    def forward(self, x):\n        # x: (B, C, H, W)\n        # Clamp to avoid numerical issues with pow\n        x_clamped = x.clamp(min=self.eps)\n        # Raise to power p, average-pool, then take the p-th root\n        pooled = F.adaptive_avg_pool2d(x_clamped.pow(self.p), 1)\n        return pooled.pow(1.0 / self.p).flatten(1)  # (B, C)\n    \n    def __repr__(self):\n        return f\"GeM(p={self.p.item():.2f})\"\n\n\nclass GAP(nn.Module):\n    \"\"\"Standard Global Average Pooling (for ablation comparison).\"\"\"\n    def forward(self, x):\n        return F.adaptive_avg_pool2d(x, 1).flatten(1)\n\n\nclass AvgGAPGMP(nn.Module):\n    \"\"\"Average of GAP and GMP (for ablation comparison).\"\"\"\n    def forward(self, x):\n        gap = F.adaptive_avg_pool2d(x, 1).flatten(1)\n        gmp = F.adaptive_max_pool2d(x, 1).flatten(1)\n        return 0.5 * (gap + gmp)\n\n\ndef build_pool(name):\n    \"\"\"Factory for pooling layers.\"\"\"\n    pools = {\n        'gem': GeM,\n        'gap': GAP,\n        'avg_gap_gmp': AvgGAPGMP,\n    }\n    assert name in pools, f\"Unknown pool: {name}. Choose from {list(pools.keys())}\"\n    return pools[name]()\n\n\n# Quick test\ndummy = torch.randn(2, 512, 4, 4)\nfor name in ['gem', 'gap', 'avg_gap_gmp']:\n    pool = build_pool(name)\n    out = pool(dummy)\n    print(f\"  {name:12s} → {out.shape}\")\nprint(\"✓ All pooling modules work\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:17:58.314014Z","iopub.execute_input":"2026-09-19T07:17:58.314427Z","iopub.status.idle":"2026-09-19T07:17:58.356144Z","shell.execute_reply.started":"2026-09-19T07:17:58.314397Z","shell.execute_reply":"2026-09-19T07:17:58.355309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 3, Cell 3: Model architecture\n\nclass BengaliModel(nn.Module):\n    \"\"\"Shared CNN trunk + output head(s).\n    \n    head_type='component' → 3 independent softmax heads (root/vowel/cons)\n    head_type='flat'      → 1 softmax head over N seen whole graphemes\n    \n    The trunk and pooling are identical in both cases.\n    Only the head differs — that's the controlled experiment.\n    \"\"\"\n    \n    def __init__(self, arch='efficientnet_b0', pretrained=True,\n                 pool_name='gem', head_type='component',\n                 n_flat_classes=None, in_chans=1):\n        super().__init__()\n        \n        self.head_type = head_type\n        \n        # 1. CNN trunk from timm (removes the default classifier head)\n        self.trunk = timm.create_model(\n            arch,\n            pretrained=pretrained,\n            in_chans=in_chans,     # 1-channel grayscale input\n            num_classes=0,          # removes the final FC\n            global_pool='',         # we'll do our own pooling\n        )\n        \n        # Get the trunk's output feature dimension\n        # timm models expose this as .num_features\n        self.feature_dim = self.trunk.num_features\n        \n        # 2. Pooling\n        self.pool = build_pool(pool_name)\n        \n        # 3. Head(s)\n        if head_type == 'component':\n            self.head_root = nn.Linear(self.feature_dim, N_ROOT)\n            self.head_vowel = nn.Linear(self.feature_dim, N_VOWEL)\n            self.head_cons = nn.Linear(self.feature_dim, N_CONS)\n        elif head_type == 'flat':\n            assert n_flat_classes is not None, \\\n                \"n_flat_classes required for flat head\"\n            self.head_flat = nn.Linear(self.feature_dim, n_flat_classes)\n        else:\n            raise ValueError(f\"Unknown head_type: {head_type}\")\n    \n    def forward(self, x):\n        \"\"\"\n        Args:\n            x: (B, 1, 128, 128) tensor\n        Returns:\n            dict of logits:\n                component → {'root': (B,168), 'vowel': (B,11), 'cons': (B,7)}\n                flat      → {'flat': (B, n_flat_classes)}\n        \"\"\"\n        # Trunk: extract features\n        features = self.trunk(x)              # (B, C, H', W')\n        \n        # Pool: collapse spatial dims\n        pooled = self.pool(features)          # (B, feature_dim)\n        \n        # Head(s)\n        if self.head_type == 'component':\n            return {\n                'root':  self.head_root(pooled),\n                'vowel': self.head_vowel(pooled),\n                'cons':  self.head_cons(pooled),\n            }\n        else:\n            return {\n                'flat': self.head_flat(pooled),\n            }\n    \n    def count_parameters(self):\n        \"\"\"Count total and trainable parameters.\"\"\"\n        total = sum(p.numel() for p in self.parameters())\n        trainable = sum(p.numel() for p in self.parameters() if p.requires_grad)\n        return total, trainable\n\n\ndef build_model(config):\n    \"\"\"Build a model from a config dict.\n    \n    config keys: arch, pretrained, pool, head_type, n_flat_classes (if flat)\n    \"\"\"\n    return BengaliModel(\n        arch=config.get('arch', 'efficientnet_b0'),\n        pretrained=config.get('pretrained', True),\n        pool_name=config.get('pool', 'gem'),\n        head_type=config.get('head_type', 'component'),\n        n_flat_classes=config.get('n_flat_classes', None),\n        in_chans=config.get('in_chans', 1),\n    )\n\n\n# ----- Test both model variants -----\nprint(\"=== Component model ===\")\nmodel_comp = build_model({\n    'arch': 'efficientnet_b0',\n    'pretrained': True,\n    'pool': 'gem',\n    'head_type': 'component',\n    'in_chans': 1,\n})\ntotal, trainable = model_comp.count_parameters()\nprint(f\"  Parameters: {total:,} total, {trainable:,} trainable\")\nprint(f\"  Feature dim: {model_comp.feature_dim}\")\nprint(f\"  GeM p: {model_comp.pool.p.item():.2f}\")\n\nprint(f\"\\n=== Flat model ===\")\nmodel_flat = build_model({\n    'arch': 'efficientnet_b0',\n    'pretrained': True,\n    'pool': 'gem',\n    'head_type': 'flat',\n    'n_flat_classes': 1036,\n    'in_chans': 1,\n})\ntotal_f, trainable_f = model_flat.count_parameters()\nprint(f\"  Parameters: {total_f:,} total, {trainable_f:,} trainable\")\nprint(f\"  Flat output: {1036} classes\")\n\n# Forward pass test\ndummy_input = torch.randn(2, 1, TARGET_SIZE, TARGET_SIZE)\nout_comp = model_comp(dummy_input)\nout_flat = model_flat(dummy_input)\n\nprint(f\"\\n=== Forward pass test ===\")\nprint(f\"  Component output:\")\nfor k, v in out_comp.items():\n    print(f\"    {k}: {v.shape}\")\nprint(f\"  Flat output:\")\nfor k, v in out_flat.items():\n    print(f\"    {k}: {v.shape}\")\n\nprint(\"\\n✓ Both model variants build and forward correctly\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:18:23.855382Z","iopub.execute_input":"2026-09-19T07:18:23.855945Z","iopub.status.idle":"2026-09-19T07:18:29.364726Z","shell.execute_reply.started":"2026-09-19T07:18:23.855915Z","shell.execute_reply":"2026-09-19T07:18:29.363941Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 3, Cell 4: Custom shallow CNN (course-honest baseline)\n\nclass ShallowCNN(nn.Module):\n    \"\"\"Simple 4-conv CNN — the 'did anything learn' baseline.\n    \n    No pretrained weights, no fancy architecture.\n    Exists to prove the multi-task setup works before we bring\n    in heavy backbones, and as the from-scratch ablation point.\n    \"\"\"\n    def __init__(self, pool_name='gem', head_type='component',\n                 n_flat_classes=None):\n        super().__init__()\n        \n        self.head_type = head_type\n        \n        # 4 conv blocks with batch norm and relu\n        self.features = nn.Sequential(\n            # Block 1: 1 → 32, 128→64\n            nn.Conv2d(1, 32, 3, padding=1), nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.Conv2d(32, 32, 3, padding=1), nn.BatchNorm2d(32), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n            nn.Dropout2d(0.1),\n            \n            # Block 2: 32 → 64, 64→32\n            nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n            nn.Dropout2d(0.1),\n            \n            # Block 3: 64 → 128, 32→16\n            nn.Conv2d(64, 128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True),\n            nn.Conv2d(128, 128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n            nn.Dropout2d(0.2),\n            \n            # Block 4: 128 → 256, 16→8\n            nn.Conv2d(128, 256, 3, padding=1), nn.BatchNorm2d(256), nn.ReLU(inplace=True),\n            nn.Conv2d(256, 256, 3, padding=1), nn.BatchNorm2d(256), nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n            nn.Dropout2d(0.2),\n        )\n        \n        self.feature_dim = 256\n        self.num_features = 256  # for compatibility\n        self.pool = build_pool(pool_name)\n        \n        if head_type == 'component':\n            self.head_root = nn.Linear(256, N_ROOT)\n            self.head_vowel = nn.Linear(256, N_VOWEL)\n            self.head_cons = nn.Linear(256, N_CONS)\n        elif head_type == 'flat':\n            assert n_flat_classes is not None\n            self.head_flat = nn.Linear(256, n_flat_classes)\n    \n    def forward(self, x):\n        features = self.features(x)     # (B, 256, 8, 8) at 128×128 input\n        pooled = self.pool(features)     # (B, 256)\n        \n        if self.head_type == 'component':\n            return {\n                'root':  self.head_root(pooled),\n                'vowel': self.head_vowel(pooled),\n                'cons':  self.head_cons(pooled),\n            }\n        else:\n            return {'flat': self.head_flat(pooled)}\n    \n    def count_parameters(self):\n        total = sum(p.numel() for p in self.parameters())\n        trainable = sum(p.numel() for p in self.parameters() if p.requires_grad)\n        return total, trainable\n\n\n# Test\nshallow = ShallowCNN(pool_name='gem', head_type='component')\ntotal, trainable = shallow.count_parameters()\nprint(f\"ShallowCNN (component):\")\nprint(f\"  Parameters: {total:,} total, {trainable:,} trainable\")\nout = shallow(torch.randn(2, 1, 128, 128))\nfor k, v in out.items():\n    print(f\"  {k}: {v.shape}\")\n\nprint(\"\\n✓ ShallowCNN builds and forwards correctly\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:18:43.336408Z","iopub.execute_input":"2026-09-19T07:18:43.337216Z","iopub.status.idle":"2026-09-19T07:18:43.425354Z","shell.execute_reply.started":"2026-09-19T07:18:43.337184Z","shell.execute_reply":"2026-09-19T07:18:43.424682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 3, Cell 5: Loss functions\n\nclass ComponentLoss(nn.Module):\n    \"\"\"Weighted sum of three cross-entropies for the component model.\n    \n    L = w_r * CE(root) + w_v * CE(vowel) + w_c * CE(cons)\n    \n    Default weights (2, 1, 1) are metric-aligned: the evaluation metric\n    weights root ×2, so the gradient signal matches what we score.\n    \"\"\"\n    def __init__(self, weights=(2.0, 1.0, 1.0), class_weights=None):\n        \"\"\"\n        Args:\n            weights: (w_root, w_vowel, w_cons) task weights\n            class_weights: dict with optional 'root', 'vowel', 'cons' tensors\n                          for class-weighted CE (inverse frequency)\n        \"\"\"\n        super().__init__()\n        self.w_r, self.w_v, self.w_c = weights\n        \n        cw = class_weights or {}\n        self.ce_root = nn.CrossEntropyLoss(weight=cw.get('root'))\n        self.ce_vowel = nn.CrossEntropyLoss(weight=cw.get('vowel'))\n        self.ce_cons = nn.CrossEntropyLoss(weight=cw.get('cons'))\n    \n    def forward(self, preds, labels):\n        \"\"\"\n        preds:  dict with 'root', 'vowel', 'cons' logits\n        labels: dict with 'root', 'vowel', 'cons' long tensors\n        \"\"\"\n        l_r = self.ce_root(preds['root'], labels['root'])\n        l_v = self.ce_vowel(preds['vowel'], labels['vowel'])\n        l_c = self.ce_cons(preds['cons'], labels['cons'])\n        \n        total = self.w_r * l_r + self.w_v * l_v + self.w_c * l_c\n        \n        return total, {'root': l_r.item(), 'vowel': l_v.item(), 'cons': l_c.item()}\n\n\nclass FocalComponentLoss(nn.Module):\n    \"\"\"Component loss with Focal Loss on the root head.\n    \n    Focal loss: FL = -(1 - p_t)^gamma * log(p_t)\n    Down-weights easy/frequent examples, concentrates gradient on\n    the hard tail. Directly targets the macro-recall objective.\n    \n    Only applied to root (168 classes, long-tailed).\n    Vowel and consonant use standard CE.\n    \"\"\"\n    def __init__(self, weights=(2.0, 1.0, 1.0), gamma=2.0):\n        super().__init__()\n        self.w_r, self.w_v, self.w_c = weights\n        self.gamma = gamma\n        self.ce_vowel = nn.CrossEntropyLoss()\n        self.ce_cons = nn.CrossEntropyLoss()\n    \n    def focal_loss(self, logits, targets):\n        ce_loss = F.cross_entropy(logits, targets, reduction='none')\n        p_t = torch.exp(-ce_loss)  # probability of correct class\n        focal = ((1 - p_t) ** self.gamma) * ce_loss\n        return focal.mean()\n    \n    def forward(self, preds, labels):\n        l_r = self.focal_loss(preds['root'], labels['root'])\n        l_v = self.ce_vowel(preds['vowel'], labels['vowel'])\n        l_c = self.ce_cons(preds['cons'], labels['cons'])\n        \n        total = self.w_r * l_r + self.w_v * l_v + self.w_c * l_c\n        return total, {'root': l_r.item(), 'vowel': l_v.item(), 'cons': l_c.item()}\n\n\nclass UncertaintyComponentLoss(nn.Module):\n    \"\"\"Learned task weighting via homoscedastic uncertainty (Kendall & Gal).\n    \n    Each task gets a learned log-sigma. Loss becomes:\n    L = sum_t [ (1 / 2*sigma_t^2) * CE_t + log(sigma_t) ]\n    \n    The model learns how much to weight each task automatically.\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        # Initialize log_sigma to 0 → sigma=1 → equal weights initially\n        self.log_sigma = nn.Parameter(torch.zeros(3))\n        self.ce_root = nn.CrossEntropyLoss()\n        self.ce_vowel = nn.CrossEntropyLoss()\n        self.ce_cons = nn.CrossEntropyLoss()\n    \n    def forward(self, preds, labels):\n        l_r = self.ce_root(preds['root'], labels['root'])\n        l_v = self.ce_vowel(preds['vowel'], labels['vowel'])\n        l_c = self.ce_cons(preds['cons'], labels['cons'])\n        \n        losses = torch.stack([l_r, l_v, l_c])\n        # precision = 1 / (2 * sigma^2) = exp(-2 * log_sigma) / 2\n        precision = torch.exp(-2 * self.log_sigma) * 0.5\n        total = (precision * losses + self.log_sigma).sum()\n        \n        return total, {\n            'root': l_r.item(), 'vowel': l_v.item(), 'cons': l_c.item(),\n            'sigma_r': torch.exp(self.log_sigma[0]).item(),\n            'sigma_v': torch.exp(self.log_sigma[1]).item(),\n            'sigma_c': torch.exp(self.log_sigma[2]).item(),\n        }\n\n\nclass FlatLoss(nn.Module):\n    \"\"\"Standard cross-entropy for the flat whole-grapheme classifier.\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.ce = nn.CrossEntropyLoss()\n    \n    def forward(self, preds, labels):\n        loss = self.ce(preds['flat'], labels['whole_idx'])\n        return loss, {'flat': loss.item()}\n\n\ndef compute_class_weights(labels_df, train_indices):\n    \"\"\"Compute inverse-frequency class weights for class-weighted CE.\n    \n    Returns dict with 'root', 'vowel', 'cons' float tensors.\n    \"\"\"\n    sub = labels_df.iloc[train_indices]\n    weights = {}\n    for col, n_classes, key in [\n        ('grapheme_root', N_ROOT, 'root'),\n        ('vowel_diacritic', N_VOWEL, 'vowel'),\n        ('consonant_diacritic', N_CONS, 'cons'),\n    ]:\n        counts = sub[col].value_counts().sort_index().values.astype(np.float32)\n        # Inverse frequency, normalized so they sum to n_classes\n        inv_freq = 1.0 / counts\n        inv_freq = inv_freq / inv_freq.sum() * n_classes\n        weights[key] = torch.from_numpy(inv_freq)\n    \n    return weights\n\n\ndef build_loss(config, labels_df=None, train_indices=None):\n    \"\"\"Build a loss function from config.\n    \n    config keys: head_type, loss_weights, imbalance, focal_gamma\n    \"\"\"\n    head_type = config.get('head_type', 'component')\n    \n    if head_type == 'flat':\n        return FlatLoss()\n    \n    loss_weights = tuple(config.get('loss_weights', [2.0, 1.0, 1.0]))\n    imbalance = config.get('imbalance', 'none')\n    \n    if imbalance == 'focal':\n        gamma = config.get('focal_gamma', 2.0)\n        return FocalComponentLoss(weights=loss_weights, gamma=gamma)\n    \n    elif imbalance == 'uncertainty':\n        return UncertaintyComponentLoss()\n    \n    elif imbalance == 'class_weighted':\n        assert labels_df is not None and train_indices is not None, \\\n            \"Need labels_df and train_indices for class_weighted\"\n        cw = compute_class_weights(labels_df, train_indices)\n        return ComponentLoss(weights=loss_weights, class_weights=cw)\n    \n    else:  # 'none'\n        return ComponentLoss(weights=loss_weights)\n\n\n# ----- Test all loss variants -----\ndummy_preds_comp = {\n    'root': torch.randn(4, N_ROOT),\n    'vowel': torch.randn(4, N_VOWEL),\n    'cons': torch.randn(4, N_CONS),\n}\ndummy_labels = {\n    'root': torch.randint(0, N_ROOT, (4,)),\n    'vowel': torch.randint(0, N_VOWEL, (4,)),\n    'cons': torch.randint(0, N_CONS, (4,)),\n    'whole_idx': torch.randint(0, 1036, (4,)),\n}\n\nprint(\"=== Loss function tests ===\")\nfor name, cfg in [\n    ('CE (2,1,1)',       {'head_type': 'component', 'loss_weights': [2,1,1], 'imbalance': 'none'}),\n    ('CE (1,1,1)',       {'head_type': 'component', 'loss_weights': [1,1,1], 'imbalance': 'none'}),\n    ('Focal (γ=2)',      {'head_type': 'component', 'imbalance': 'focal', 'focal_gamma': 2.0}),\n    ('Uncertainty',      {'head_type': 'component', 'imbalance': 'uncertainty'}),\n    ('Flat CE',          {'head_type': 'flat'}),\n]:\n    loss_fn = build_loss(cfg)\n    if cfg['head_type'] == 'flat':\n        preds = {'flat': torch.randn(4, 1036)}\n        loss, details = loss_fn(preds, dummy_labels)\n    else:\n        loss, details = loss_fn(dummy_preds_comp, dummy_labels)\n    \n    print(f\"  {name:20s} → loss={loss.item():.4f}  details={details}\")\n\nprint(\"\\n✓ All loss functions work\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:19:01.665598Z","iopub.execute_input":"2026-09-19T07:19:01.666439Z","iopub.status.idle":"2026-09-19T07:19:01.708581Z","shell.execute_reply.started":"2026-09-19T07:19:01.666405Z","shell.execute_reply":"2026-09-19T07:19:01.707959Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 3, Cell 6: Evaluation metrics\n\ndef hierarchical_macro_recall(all_preds, all_labels):\n    \"\"\"Compute the competition metric: hierarchical macro-averaged recall.\n    \n    Score = (2 * R_root + R_vowel + R_cons) / 4\n    \n    Macro-recall per component, root weighted ×2.\n    \n    Args:\n        all_preds:  dict with 'root', 'vowel', 'cons' → numpy int arrays\n        all_labels: dict with 'root', 'vowel', 'cons' → numpy int arrays\n    \n    Returns:\n        dict with 'score', 'root', 'vowel', 'cons' recall values\n    \"\"\"\n    results = {}\n    for key in ['root', 'vowel', 'cons']:\n        results[key] = recall_score(\n            all_labels[key], all_preds[key],\n            average='macro', zero_division=0\n        )\n    \n    results['score'] = (2 * results['root'] + results['vowel'] + results['cons']) / 4\n    return results\n\n\ndef whole_grapheme_accuracy(all_preds, all_labels):\n    \"\"\"Fraction of samples where ALL THREE components are correct.\"\"\"\n    correct = (\n        (all_preds['root'] == all_labels['root']) &\n        (all_preds['vowel'] == all_labels['vowel']) &\n        (all_preds['cons'] == all_labels['cons'])\n    )\n    return correct.mean()\n\n\n@torch.no_grad()\ndef evaluate_component_model(model, dataloader, device):\n    \"\"\"Run a component model on a dataloader and return metrics.\n    \n    Returns:\n        metrics: dict with 'score', 'root', 'vowel', 'cons', 'whole_acc'\n        all_preds: dict with 'root', 'vowel', 'cons' numpy arrays\n        all_labels: dict with 'root', 'vowel', 'cons' numpy arrays\n    \"\"\"\n    model.eval()\n    \n    preds = {'root': [], 'vowel': [], 'cons': []}\n    labels = {'root': [], 'vowel': [], 'cons': []}\n    \n    for batch_img, batch_labels in dataloader:\n        batch_img = batch_img.to(device)\n        out = model(batch_img)\n        \n        for key in ['root', 'vowel', 'cons']:\n            preds[key].append(out[key].argmax(dim=1).cpu().numpy())\n            labels[key].append(batch_labels[key].numpy())\n    \n    # Concatenate\n    for key in ['root', 'vowel', 'cons']:\n        preds[key] = np.concatenate(preds[key])\n        labels[key] = np.concatenate(labels[key])\n    \n    metrics = hierarchical_macro_recall(preds, labels)\n    metrics['whole_acc'] = float(whole_grapheme_accuracy(preds, labels))\n    \n    return metrics, preds, labels\n\n\n@torch.no_grad()\ndef evaluate_flat_model(model, dataloader, idx_to_components, device):\n    \"\"\"Run a flat model on a dataloader, decompose predictions into\n    components, and return the same metrics as evaluate_component_model.\n    \n    This is how we score the flat baseline on the same yardstick.\n    \n    Args:\n        idx_to_components: dict mapping flat class index → \n                          {'root': int, 'vowel': int, 'cons': int}\n    \"\"\"\n    model.eval()\n    \n    preds = {'root': [], 'vowel': [], 'cons': []}\n    labels = {'root': [], 'vowel': [], 'cons': []}\n    \n    for batch_img, batch_labels in dataloader:\n        batch_img = batch_img.to(device)\n        out = model(batch_img)\n        \n        # Get flat predictions\n        flat_preds = out['flat'].argmax(dim=1).cpu().numpy()  # (B,)\n        \n        # Decompose each flat prediction into components\n        for pred_idx in flat_preds:\n            if pred_idx in idx_to_components:\n                comp = idx_to_components[pred_idx]\n                preds['root'].append(comp['root'])\n                preds['vowel'].append(comp['vowel'])\n                preds['cons'].append(comp['cons'])\n            else:\n                # Should not happen (all flat preds are within seen range)\n                preds['root'].append(-1)\n                preds['vowel'].append(-1)\n                preds['cons'].append(-1)\n        \n        for key in ['root', 'vowel', 'cons']:\n            labels[key].append(batch_labels[key].numpy())\n    \n    # Concatenate\n    for key in ['root', 'vowel', 'cons']:\n        preds[key] = np.array(preds[key])\n        labels[key] = np.concatenate(labels[key])\n    \n    metrics = hierarchical_macro_recall(preds, labels)\n    metrics['whole_acc'] = float(whole_grapheme_accuracy(preds, labels))\n    \n    return metrics, preds, labels\n\n\n# ----- Test with dummy predictions -----\nnp.random.seed(42)\ndummy_p = {\n    'root': np.random.randint(0, N_ROOT, 100),\n    'vowel': np.random.randint(0, N_VOWEL, 100),\n    'cons': np.random.randint(0, N_CONS, 100),\n}\ndummy_l = {\n    'root': np.random.randint(0, N_ROOT, 100),\n    'vowel': np.random.randint(0, N_VOWEL, 100),\n    'cons': np.random.randint(0, N_CONS, 100),\n}\n\nmetrics = hierarchical_macro_recall(dummy_p, dummy_l)\nwhole_acc = whole_grapheme_accuracy(dummy_p, dummy_l)\n\nprint(\"=== Metric test (random predictions) ===\")\nprint(f\"  Score:    {metrics['score']:.4f}\")\nprint(f\"  Root R:   {metrics['root']:.4f}\")\nprint(f\"  Vowel R:  {metrics['vowel']:.4f}\")\nprint(f\"  Cons R:   {metrics['cons']:.4f}\")\nprint(f\"  Whole acc: {whole_acc:.4f}\")\nprint(f\"  (Random on {N_ROOT} classes ≈ {1/N_ROOT:.4f} — sanity check)\")\n\nassert 0 <= metrics['score'] <= 1.0\nassert 0 <= whole_acc <= 1.0\nprint(\"\\n✓ Metrics work correctly\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:19:16.09246Z","iopub.execute_input":"2026-09-19T07:19:16.093199Z","iopub.status.idle":"2026-09-19T07:19:16.118113Z","shell.execute_reply.started":"2026-09-19T07:19:16.093163Z","shell.execute_reply":"2026-09-19T07:19:16.117259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 3, Cell 7: Smoke test — forward/backward on one real batch,\n# then verify loss decreases over 5 steps\n\n# Reload what we need from Phase 2\nimport json\n\nWORK_DIR = Path('/kaggle/working')\nCACHE_DIR = WORK_DIR / 'cache'\nSPLIT_DIR = WORK_DIR / 'splits'\n\nimages = np.load(CACHE_DIR / 'images_128.npy', mmap_mode='r')\nlabels_df = pd.read_parquet(CACHE_DIR / 'labels.parquet')\n\nwith open(SPLIT_DIR / 'split_holdout_20pct_seed42.json', 'r') as f:\n    saved = json.load(f)\nsplit_main = saved['indices']\n\n# Build whole_to_idx (need this from Phase 2)\ntrain_graphemes = sorted(labels_df.iloc[split_main['train']]['grapheme'].unique())\nwhole_to_idx = {g: i for i, g in enumerate(train_graphemes)}\n\n# We need to import the Dataset and transforms from Phase 2\n# (In Kaggle they're still in memory; if restarted, re-run Phase 2 cells 1,6,7)\n\n# Build a small train loader for the smoke test\nsmoke_ds = BengaliDataset(\n    images_mmap=images,\n    labels_df=labels_df,\n    indices=split_main['train'][:256],  # just 256 samples\n    transform=get_train_transform(),\n    whole_to_idx=whole_to_idx,\n)\nsmoke_dl = DataLoader(smoke_ds, batch_size=32, shuffle=True, num_workers=0)\n\nprint(\"=== Smoke test: component model ===\")\nmodel = build_model({\n    'arch': 'efficientnet_b0', 'pretrained': True,\n    'pool': 'gem', 'head_type': 'component', 'in_chans': 1,\n}).to(device)\n\nloss_fn = ComponentLoss(weights=(2.0, 1.0, 1.0))\noptimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.01)\nscaler = torch.cuda.amp.GradScaler()\n\nmodel.train()\nlosses = []\n\nfor step in range(5):\n    batch_img, batch_labels = next(iter(smoke_dl))\n    batch_img = batch_img.to(device)\n    labels_gpu = {k: v.to(device) for k, v in batch_labels.items()}\n    \n    optimizer.zero_grad()\n    \n    with torch.cuda.amp.autocast():\n        preds = model(batch_img)\n        loss, details = loss_fn(preds, labels_gpu)\n    \n    scaler.scale(loss).backward()\n    scaler.step(optimizer)\n    scaler.update()\n    \n    losses.append(loss.item())\n    print(f\"  Step {step}: loss={loss.item():.4f}  \"\n          f\"root={details['root']:.3f} vowel={details['vowel']:.3f} cons={details['cons']:.3f}\")\n\nprint(f\"\\n  Loss trend: {' → '.join(f'{l:.3f}' for l in losses)}\")\nif losses[-1] < losses[0]:\n    print(f\"  ✓ Loss decreased ({losses[0]:.3f} → {losses[-1]:.3f}) — learning is happening\")\nelse:\n    print(f\"  ⚠ Loss didn't decrease — may need more steps or LR adjustment\")\n\n# Check GeM p moved\nprint(f\"  GeM p after 5 steps: {model.pool.p.item():.4f} (started at 3.0)\")\n\n# Check gradients exist on all heads\nfor name, param in model.named_parameters():\n    if 'head' in name and param.grad is not None:\n        grad_norm = param.grad.norm().item()\n        if grad_norm == 0:\n            print(f\"  ⚠ Zero gradient on {name}\")\n        break\nprint(f\"  ✓ Gradients flowing to heads\")\n\n\n# ----- Same test for flat model -----\nprint(f\"\\n=== Smoke test: flat model ===\")\nmodel_f = build_model({\n    'arch': 'efficientnet_b0', 'pretrained': True,\n    'pool': 'gem', 'head_type': 'flat',\n    'n_flat_classes': len(whole_to_idx), 'in_chans': 1,\n}).to(device)\n\nloss_fn_f = FlatLoss()\noptimizer_f = torch.optim.AdamW(model_f.parameters(), lr=3e-4, weight_decay=0.01)\nscaler_f = torch.cuda.amp.GradScaler()\n\nmodel_f.train()\nlosses_f = []\n\nfor step in range(5):\n    batch_img, batch_labels = next(iter(smoke_dl))\n    batch_img = batch_img.to(device)\n    labels_gpu = {k: v.to(device) for k, v in batch_labels.items()}\n    \n    optimizer_f.zero_grad()\n    \n    with torch.cuda.amp.autocast():\n        preds = model_f(batch_img)\n        loss, details = loss_fn_f(preds, labels_gpu)\n    \n    scaler_f.scale(loss).backward()\n    scaler_f.step(optimizer_f)\n    scaler_f.update()\n    \n    losses_f.append(loss.item())\n    print(f\"  Step {step}: loss={loss.item():.4f}\")\n\nprint(f\"\\n  Loss trend: {' → '.join(f'{l:.3f}' for l in losses_f)}\")\nif losses_f[-1] < losses_f[0]:\n    print(f\"  ✓ Loss decreased ({losses_f[0]:.3f} → {losses_f[-1]:.3f})\")\nelse:\n    print(f\"  ⚠ Loss didn't decrease — check label alignment\")\n\n# Cleanup\ndel model, model_f, optimizer, optimizer_f, scaler, scaler_f\ntorch.cuda.empty_cache()\ngc.collect()\n\nprint(f\"\\n✓ Phase 3 complete — models, losses, and metrics verified\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:19:33.729372Z","iopub.execute_input":"2026-09-19T07:19:33.730595Z","iopub.status.idle":"2026-09-19T07:19:45.065355Z","shell.execute_reply.started":"2026-09-19T07:19:33.730557Z","shell.execute_reply":"2026-09-19T07:19:45.064705Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Phase 4 — The Training Loop + Main Experiment","metadata":{}},{"cell_type":"code","source":"# Phase 4, Cell 1: Training utilities\n\nimport os\nimport gc\nimport csv\nimport json\nimport time\nimport copy\nfrom pathlib import Path\nfrom collections import defaultdict\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\n\n# Paths\nWORK_DIR = Path('/kaggle/working')\nCACHE_DIR = WORK_DIR / 'cache'\nSPLIT_DIR = WORK_DIR / 'splits'\nFIG_DIR = WORK_DIR / 'figures'\nCKPT_DIR = WORK_DIR / 'checkpoints'\nLOG_DIR = WORK_DIR / 'logs'\nCKPT_DIR.mkdir(parents=True, exist_ok=True)\nLOG_DIR.mkdir(parents=True, exist_ok=True)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n\ndef set_seed(seed):\n    \"\"\"Set all random seeds for reproducibility.\"\"\"\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\n\ndef build_scheduler(optimizer, warmup_epochs, total_epochs, steps_per_epoch):\n    \"\"\"Linear warmup → cosine decay to 0.\"\"\"\n    warmup_steps = warmup_epochs * steps_per_epoch\n    cosine_steps = (total_epochs - warmup_epochs) * steps_per_epoch\n\n    warmup = LinearLR(\n        optimizer,\n        start_factor=0.01,\n        end_factor=1.0,\n        total_iters=warmup_steps,\n    )\n    cosine = CosineAnnealingLR(\n        optimizer,\n        T_max=cosine_steps,\n        eta_min=1e-7,\n    )\n    scheduler = SequentialLR(\n        optimizer,\n        schedulers=[warmup, cosine],\n        milestones=[warmup_steps],\n    )\n    return scheduler\n\n\nclass CSVLogger:\n    \"\"\"Append one row per epoch to a CSV file.\"\"\"\n    def __init__(self, filepath, fieldnames):\n        self.filepath = filepath\n        self.fieldnames = fieldnames\n        with open(filepath, 'w', newline='') as f:\n            writer = csv.DictWriter(f, fieldnames=fieldnames)\n            writer.writeheader()\n\n    def log(self, row):\n        with open(self.filepath, 'a', newline='') as f:\n            writer = csv.DictWriter(f, fieldnames=self.fieldnames)\n            writer.writerow(row)\n\n\nprint(f\"✓ Training utilities ready\")\nprint(f\"  Device: {device}\")\nprint(f\"  Checkpoints: {CKPT_DIR}\")\nprint(f\"  Logs: {LOG_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:26:32.299477Z","iopub.execute_input":"2026-09-19T07:26:32.300105Z","iopub.status.idle":"2026-09-19T07:26:32.311928Z","shell.execute_reply.started":"2026-09-19T07:26:32.300075Z","shell.execute_reply":"2026-09-19T07:26:32.311113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 4, Cell 2: The training loop\n\ndef train_one_epoch(model, dataloader, loss_fn, optimizer, scheduler, scaler, device,\n                    use_cutmix=False, use_mixup=False, cutmix_alpha=1.0, mixup_alpha=0.4):\n    \"\"\"Train for one epoch. Returns average loss.\"\"\"\n    model.train()\n    total_loss = 0.0\n    n_batches = 0\n\n    for batch_img, batch_labels in dataloader:\n        batch_img = batch_img.to(device, non_blocking=True)\n        labels_gpu = {k: v.to(device, non_blocking=True) for k, v in batch_labels.items()}\n\n        # Optional CutMix / MixUp\n        mixed = False\n        if use_cutmix and np.random.random() < 0.5:\n            batch_img, labels_a, labels_b, lam = cutmix_batch(batch_img, labels_gpu, cutmix_alpha)\n            mixed = True\n        elif use_mixup and np.random.random() < 0.5:\n            batch_img, labels_a, labels_b, lam = mixup_batch(batch_img, labels_gpu, mixup_alpha)\n            mixed = True\n\n        optimizer.zero_grad()\n\n        with torch.amp.autocast('cuda'):\n            preds = model(batch_img)\n            if mixed:\n                loss = mix_criterion(loss_fn, preds, labels_a, labels_b, lam)\n                if isinstance(loss, tuple):\n                    loss = loss[0]  # mix_criterion may return (loss, details)\n            else:\n                loss_out = loss_fn(preds, labels_gpu)\n                loss = loss_out[0] if isinstance(loss_out, tuple) else loss_out\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n\n        total_loss += loss.item()\n        n_batches += 1\n\n    return total_loss / max(n_batches, 1)\n\n\ndef train_model(config):\n    \"\"\"Full training pipeline for one run.\n\n    Args:\n        config: dict with all hyperparameters (see below for keys)\n\n    Returns:\n        best_metrics: dict with best validation metrics\n        history: list of dicts (one per epoch)\n    \"\"\"\n    run_name = config['run_name']\n    seed = config['seed']\n    set_seed(seed)\n\n    print(f\"\\n{'='*60}\")\n    print(f\"  RUN: {run_name}\")\n    print(f\"  Seed: {seed}  |  Arch: {config['arch']}  |  Head: {config['head_type']}\")\n    print(f\"{'='*60}\")\n\n    # ---- Data ----\n    split_path = SPLIT_DIR / config['split_file']\n    with open(split_path, 'r') as f:\n        saved = json.load(f)\n    split = saved['indices']\n\n    images = np.load(CACHE_DIR / 'images_128.npy', mmap_mode='r')\n    labels_df = pd.read_parquet(CACHE_DIR / 'labels.parquet')\n\n    # Build whole_to_idx from this split's training set\n    train_graphemes = sorted(labels_df.iloc[split['train']]['grapheme'].unique())\n    w2i = {g: i for i, g in enumerate(train_graphemes)}\n\n    # Build idx_to_components for flat model evaluation\n    i2c = {}\n    for g in train_graphemes:\n        row = labels_df[labels_df['grapheme'] == g].iloc[0]\n        i2c[w2i[g]] = {\n            'root': int(row['grapheme_root']),\n            'vowel': int(row['vowel_diacritic']),\n            'cons': int(row['consonant_diacritic']),\n        }\n\n    train_ds = BengaliDataset(images, labels_df, split['train'],\n                              transform=get_train_transform(), whole_to_idx=w2i)\n    val_ds = BengaliDataset(images, labels_df, split['val'],\n                            transform=get_val_transform(), whole_to_idx=w2i)\n    seen_ds = BengaliDataset(images, labels_df, split['seen_test'],\n                             transform=get_val_transform(), whole_to_idx=w2i)\n    unseen_ds = BengaliDataset(images, labels_df, split['unseen_test'],\n                               transform=get_val_transform(), whole_to_idx=w2i)\n\n    bs = config.get('batch_size', 64)\n    nw = config.get('num_workers', 2)\n    train_dl = DataLoader(train_ds, batch_size=bs, shuffle=True,\n                          num_workers=nw, pin_memory=True, drop_last=True)\n    val_dl = DataLoader(val_ds, batch_size=bs * 2, shuffle=False,\n                        num_workers=nw, pin_memory=True)\n    seen_dl = DataLoader(seen_ds, batch_size=bs * 2, shuffle=False,\n                         num_workers=nw, pin_memory=True)\n    unseen_dl = DataLoader(unseen_ds, batch_size=bs * 2, shuffle=False,\n                           num_workers=nw, pin_memory=True)\n\n    # ---- Model ----\n    model_cfg = {\n        'arch': config['arch'],\n        'pretrained': config.get('pretrained', True),\n        'pool': config.get('pool', 'gem'),\n        'head_type': config['head_type'],\n        'n_flat_classes': len(w2i) if config['head_type'] == 'flat' else None,\n        'in_chans': 1,\n    }\n\n    if config['arch'] == 'shallow_cnn':\n        model = ShallowCNN(\n            pool_name=config.get('pool', 'gem'),\n            head_type=config['head_type'],\n            n_flat_classes=len(w2i) if config['head_type'] == 'flat' else None,\n        ).to(device)\n    else:\n        model = build_model(model_cfg).to(device)\n\n    total_params, train_params = model.count_parameters()\n    print(f\"  Parameters: {total_params:,} total, {train_params:,} trainable\")\n\n    # ---- Loss ----\n    loss_fn = build_loss(config, labels_df, split['train'])\n    if hasattr(loss_fn, 'log_sigma'):\n        # Uncertainty loss has learnable params — add to optimizer\n        loss_fn = loss_fn.to(device)\n\n    # ---- Optimizer ----\n    lr = config.get('lr', 3e-4)\n    wd = config.get('weight_decay', 0.01)\n    params = list(model.parameters())\n    if hasattr(loss_fn, 'parameters'):\n        params += list(loss_fn.parameters())\n    optimizer = torch.optim.AdamW(params, lr=lr, weight_decay=wd)\n\n    # ---- Scheduler ----\n    epochs = config.get('epochs', 40)\n    warmup_epochs = config.get('warmup_epochs', 3)\n    steps_per_epoch = len(train_dl)\n    scheduler = build_scheduler(optimizer, warmup_epochs, epochs, steps_per_epoch)\n\n    # ---- AMP ----\n    scaler = torch.amp.GradScaler('cuda')\n\n    # ---- Logging ----\n    log_fields = ['epoch', 'lr', 'train_loss', 'val_score',\n                  'val_root', 'val_vowel', 'val_cons', 'val_whole_acc']\n    logger = CSVLogger(LOG_DIR / f'{run_name}.csv', log_fields)\n\n    # ---- Training loop ----\n    best_val_score = -1.0\n    best_epoch = -1\n    patience = config.get('patience', 10)\n    patience_counter = 0\n    history = []\n\n    use_cutmix = config.get('use_cutmix', False)\n    use_mixup = config.get('use_mixup', False)\n\n    t_start = time.time()\n\n    for epoch in range(1, epochs + 1):\n        t_epoch = time.time()\n\n        # Train\n        train_loss = train_one_epoch(\n            model, train_dl, loss_fn, optimizer, scheduler, scaler, device,\n            use_cutmix=use_cutmix, use_mixup=use_mixup,\n            cutmix_alpha=config.get('cutmix_alpha', 1.0),\n            mixup_alpha=config.get('mixup_alpha', 0.4),\n        )\n\n        # Validate\n        if config['head_type'] == 'component':\n            val_metrics, _, _ = evaluate_component_model(model, val_dl, device)\n        else:\n            val_metrics, _, _ = evaluate_flat_model(model, val_dl, i2c, device)\n\n        current_lr = optimizer.param_groups[0]['lr']\n        epoch_time = time.time() - t_epoch\n\n        # Log\n        row = {\n            'epoch': epoch,\n            'lr': f\"{current_lr:.2e}\",\n            'train_loss': f\"{train_loss:.4f}\",\n            'val_score': f\"{val_metrics['score']:.4f}\",\n            'val_root': f\"{val_metrics['root']:.4f}\",\n            'val_vowel': f\"{val_metrics['vowel']:.4f}\",\n            'val_cons': f\"{val_metrics['cons']:.4f}\",\n            'val_whole_acc': f\"{val_metrics['whole_acc']:.4f}\",\n        }\n        logger.log(row)\n        history.append({**row, 'epoch_time': epoch_time})\n\n        # Print progress\n        if epoch <= 3 or epoch % 5 == 0 or epoch == epochs:\n            print(f\"  Ep {epoch:3d}/{epochs} | loss={train_loss:.4f} | \"\n                  f\"val={val_metrics['score']:.4f} (r={val_metrics['root']:.3f} \"\n                  f\"v={val_metrics['vowel']:.3f} c={val_metrics['cons']:.3f}) | \"\n                  f\"lr={current_lr:.1e} | {epoch_time:.0f}s\")\n\n        # Checkpoint best\n        if val_metrics['score'] > best_val_score:\n            best_val_score = val_metrics['score']\n            best_epoch = epoch\n            patience_counter = 0\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'val_score': best_val_score,\n                'config': config,\n            }, CKPT_DIR / f'{run_name}_best.pt')\n        else:\n            patience_counter += 1\n\n        # Early stopping\n        if patience_counter >= patience:\n            print(f\"  Early stopping at epoch {epoch} (patience={patience})\")\n            break\n\n    total_time = time.time() - t_start\n    print(f\"\\n  Training complete in {total_time / 60:.1f} min\")\n    print(f\"  Best val score: {best_val_score:.4f} at epoch {best_epoch}\")\n\n    # ---- Final evaluation on Seen + Unseen test sets ----\n    print(f\"\\n  Loading best checkpoint (epoch {best_epoch})...\")\n    ckpt = torch.load(CKPT_DIR / f'{run_name}_best.pt', map_location=device,\n                      weights_only=False)\n    model.load_state_dict(ckpt['model_state_dict'])\n\n    print(f\"  Evaluating on Seen test...\")\n    if config['head_type'] == 'component':\n        seen_metrics, _, _ = evaluate_component_model(model, seen_dl, device)\n    else:\n        seen_metrics, _, _ = evaluate_flat_model(model, seen_dl, i2c, device)\n\n    print(f\"  Evaluating on Unseen test...\")\n    if config['head_type'] == 'component':\n        unseen_metrics, _, _ = evaluate_component_model(model, unseen_dl, device)\n    else:\n        unseen_metrics, _, _ = evaluate_flat_model(model, unseen_dl, i2c, device)\n\n    gap = seen_metrics['score'] - unseen_metrics['score']\n\n    print(f\"\\n  ┌───────────────────────────────────────────────────┐\")\n    print(f\"  │ RESULTS: {run_name:40s}│\")\n    print(f\"  ├───────────┬──────────┬──────────┬────────────────┤\")\n    print(f\"  │           │  Score   │  Root    │  Whole Acc     │\")\n    print(f\"  ├───────────┼──────────┼──────────┼────────────────┤\")\n    print(f\"  │ Seen      │  {seen_metrics['score']:.4f}  │  {seen_metrics['root']:.4f}  │  {seen_metrics['whole_acc']:.4f}          │\")\n    print(f\"  │ Unseen    │  {unseen_metrics['score']:.4f}  │  {unseen_metrics['root']:.4f}  │  {unseen_metrics['whole_acc']:.4f}          │\")\n    print(f\"  │ Gap       │  {gap:.4f}  │          │                │\")\n    print(f\"  └───────────┴──────────┴──────────┴────────────────┘\")\n\n    final_results = {\n        'run_name': run_name,\n        'config': config,\n        'best_epoch': best_epoch,\n        'best_val_score': best_val_score,\n        'seen': seen_metrics,\n        'unseen': unseen_metrics,\n        'gap': gap,\n        'total_time_min': total_time / 60,\n    }\n\n    # Save results JSON\n    with open(LOG_DIR / f'{run_name}_results.json', 'w') as f:\n        json.dump(final_results, f, indent=2)\n\n    # Cleanup\n    del model, optimizer, scaler, train_dl, val_dl, seen_dl, unseen_dl\n    del images\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    return final_results, history\n\n\nprint(\"✓ train_model() defined — ready for experiments\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:26:45.128533Z","iopub.execute_input":"2026-09-19T07:26:45.129174Z","iopub.status.idle":"2026-09-19T07:26:45.158116Z","shell.execute_reply.started":"2026-09-19T07:26:45.129143Z","shell.execute_reply":"2026-09-19T07:26:45.157264Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 4, Cell 3: Updated mix_criterion that works with both loss types\n\ndef mix_criterion(loss_fn, preds, labels_a, labels_b, lam):\n    \"\"\"Compute mixed loss for CutMix/MixUp.\n\n    Works with both ComponentLoss (returns tuple) and FlatLoss.\n    \"\"\"\n    out_a = loss_fn(preds, labels_a)\n    out_b = loss_fn(preds, labels_b)\n\n    loss_a = out_a[0] if isinstance(out_a, tuple) else out_a\n    loss_b = out_b[0] if isinstance(out_b, tuple) else out_b\n\n    return lam * loss_a + (1.0 - lam) * loss_b\n\nprint(\"✓ mix_criterion updated\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:27:02.144086Z","iopub.execute_input":"2026-09-19T07:27:02.144628Z","iopub.status.idle":"2026-09-19T07:27:02.15003Z","shell.execute_reply.started":"2026-09-19T07:27:02.144598Z","shell.execute_reply":"2026-09-19T07:27:02.149289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 4, Cell 4: Reference run — component model, seed 42\n\nreference_config = {\n    'run_name': 'effb0_component_seed42',\n    'seed': 42,\n    'arch': 'efficientnet_b0',\n    'pretrained': True,\n    'pool': 'gem',\n    'head_type': 'component',\n    'split_file': 'split_holdout_20pct_seed42.json',\n\n    # Loss\n    'loss_weights': [2.0, 1.0, 1.0],\n    'imbalance': 'none',\n\n    # Training\n    'batch_size': 64,\n    'lr': 3e-4,\n    'weight_decay': 0.01,\n    'epochs': 40,\n    'warmup_epochs': 3,\n    'patience': 10,\n    'num_workers': 2,\n\n    # Augmentation\n    'use_cutmix': False,\n    'use_mixup': False,\n}\n\nref_results, ref_history = train_model(reference_config)\n\n# Quick sanity: did we reach a reasonable seen score?\nseen_score = ref_results['seen']['score']\nif seen_score > 0.85:\n    print(f\"\\n✓ Reference run looks good (seen score: {seen_score:.4f})\")\nelse:\n    print(f\"\\n⚠ Seen score is {seen_score:.4f} — lower than expected. Check training logs.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T07:27:13.965Z","iopub.execute_input":"2026-09-19T07:27:13.965653Z","iopub.status.idle":"2026-09-19T09:02:58.789507Z","shell.execute_reply.started":"2026-09-19T07:27:13.965614Z","shell.execute_reply":"2026-09-19T09:02:58.788829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 4, Cell 5: Plot training curve for the reference run\n\nref_log = pd.read_csv(LOG_DIR / 'effb0_component_seed42.csv')\n\nfig, axes = plt.subplots(1, 3, figsize=(18, 4.5))\n\n# Panel 1: Training loss\naxes[0].plot(ref_log['epoch'], ref_log['train_loss'].astype(float), color='#C0392B', linewidth=2)\naxes[0].set_title('Training loss', fontsize=13)\naxes[0].set_xlabel('Epoch')\naxes[0].set_ylabel('Loss')\naxes[0].grid(alpha=0.3)\n\n# Panel 2: Validation score\naxes[1].plot(ref_log['epoch'], ref_log['val_score'].astype(float),\n             color='#1F6F6B', linewidth=2, label='Overall score')\naxes[1].plot(ref_log['epoch'], ref_log['val_root'].astype(float),\n             color='#C0392B', linewidth=1.5, alpha=0.7, linestyle='--', label='Root')\naxes[1].plot(ref_log['epoch'], ref_log['val_vowel'].astype(float),\n             color='#5A6478', linewidth=1.5, alpha=0.7, linestyle='--', label='Vowel')\naxes[1].plot(ref_log['epoch'], ref_log['val_cons'].astype(float),\n             color='#9AA3B6', linewidth=1.5, alpha=0.7, linestyle='--', label='Cons')\naxes[1].set_title('Validation macro-recall', fontsize=13)\naxes[1].set_xlabel('Epoch')\naxes[1].set_ylabel('Macro-recall')\naxes[1].legend(fontsize=10)\naxes[1].grid(alpha=0.3)\n\n# Panel 3: Learning rate\naxes[2].plot(ref_log['epoch'], ref_log['lr'].astype(float), color='#16203A', linewidth=2)\naxes[2].set_title('Learning rate schedule', fontsize=13)\naxes[2].set_xlabel('Epoch')\naxes[2].set_ylabel('LR')\naxes[2].set_yscale('log')\naxes[2].grid(alpha=0.3)\n\nplt.suptitle(f\"Reference run: EfficientNet-B0 + GeM + component (seed 42)\", fontsize=14, y=1.02)\nplt.tight_layout()\nplt.savefig(FIG_DIR / 'reference_training_curve.png', dpi=140, bbox_inches='tight')\nplt.show()\n\n# Print final numbers\nprint(f\"\\nReference run summary:\")\nprint(f\"  Best val score: {ref_results['best_val_score']:.4f} (epoch {ref_results['best_epoch']})\")\nprint(f\"  Seen test:      {ref_results['seen']['score']:.4f}\")\nprint(f\"  Unseen test:    {ref_results['unseen']['score']:.4f}\")\nprint(f\"  Gap:            {ref_results['gap']:.4f}\")\nprint(f\"  Training time:  {ref_results['total_time_min']:.1f} min\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T09:06:11.555778Z","iopub.execute_input":"2026-09-19T09:06:11.556517Z","iopub.status.idle":"2026-09-19T09:06:12.775415Z","shell.execute_reply.started":"2026-09-19T09:06:11.556482Z","shell.execute_reply":"2026-09-19T09:06:12.774755Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 4, Cell 6: Main experiment A1 — flat vs component, 3 seeds\n\nall_results = []\n\n# If reference run already completed, include it\nif 'ref_results' in dir() and ref_results is not None:\n    all_results.append(ref_results)\n\nexperiment_configs = [\n    # Component model — seeds 43, 44 (seed 42 already done as reference)\n    {\n        'run_name': 'effb0_component_seed43',\n        'seed': 43,\n        'arch': 'efficientnet_b0', 'pretrained': True,\n        'pool': 'gem', 'head_type': 'component',\n        'split_file': 'split_holdout_20pct_seed42.json',\n        'loss_weights': [2.0, 1.0, 1.0], 'imbalance': 'none',\n        'batch_size': 64, 'lr': 3e-4, 'weight_decay': 0.01,\n        'epochs': 40, 'warmup_epochs': 3, 'patience': 10,\n        'num_workers': 2, 'use_cutmix': False, 'use_mixup': False,\n    },\n    {\n        'run_name': 'effb0_component_seed44',\n        'seed': 44,\n        'arch': 'efficientnet_b0', 'pretrained': True,\n        'pool': 'gem', 'head_type': 'component',\n        'split_file': 'split_holdout_20pct_seed42.json',\n        'loss_weights': [2.0, 1.0, 1.0], 'imbalance': 'none',\n        'batch_size': 64, 'lr': 3e-4, 'weight_decay': 0.01,\n        'epochs': 40, 'warmup_epochs': 3, 'patience': 10,\n        'num_workers': 2, 'use_cutmix': False, 'use_mixup': False,\n    },\n    # Flat model — seeds 42, 43, 44\n    {\n        'run_name': 'effb0_flat_seed42',\n        'seed': 42,\n        'arch': 'efficientnet_b0', 'pretrained': True,\n        'pool': 'gem', 'head_type': 'flat',\n        'split_file': 'split_holdout_20pct_seed42.json',\n        'loss_weights': [2.0, 1.0, 1.0], 'imbalance': 'none',\n        'batch_size': 64, 'lr': 3e-4, 'weight_decay': 0.01,\n        'epochs': 40, 'warmup_epochs': 3, 'patience': 10,\n        'num_workers': 2, 'use_cutmix': False, 'use_mixup': False,\n    },\n    {\n        'run_name': 'effb0_flat_seed43',\n        'seed': 43,\n        'arch': 'efficientnet_b0', 'pretrained': True,\n        'pool': 'gem', 'head_type': 'flat',\n        'split_file': 'split_holdout_20pct_seed42.json',\n        'loss_weights': [2.0, 1.0, 1.0], 'imbalance': 'none',\n        'batch_size': 64, 'lr': 3e-4, 'weight_decay': 0.01,\n        'epochs': 40, 'warmup_epochs': 3, 'patience': 10,\n        'num_workers': 2, 'use_cutmix': False, 'use_mixup': False,\n    },\n    {\n        'run_name': 'effb0_flat_seed44',\n        'seed': 44,\n        'arch': 'efficientnet_b0', 'pretrained': True,\n        'pool': 'gem', 'head_type': 'flat',\n        'split_file': 'split_holdout_20pct_seed42.json',\n        'loss_weights': [2.0, 1.0, 1.0], 'imbalance': 'none',\n        'batch_size': 64, 'lr': 3e-4, 'weight_decay': 0.01,\n        'epochs': 40, 'warmup_epochs': 3, 'patience': 10,\n        'num_workers': 2, 'use_cutmix': False, 'use_mixup': False,\n    },\n]\n\nfor cfg in experiment_configs:\n    # Skip if already done (checkpoint exists)\n    ckpt_path = CKPT_DIR / f\"{cfg['run_name']}_best.pt\"\n    results_path = LOG_DIR / f\"{cfg['run_name']}_results.json\"\n\n    if results_path.exists():\n        print(f\"\\n⏭ {cfg['run_name']} already done — loading results\")\n        with open(results_path, 'r') as f:\n            results = json.load(f)\n        all_results.append(results)\n        continue\n\n    results, history = train_model(cfg)\n    all_results.append(results)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"  ALL {len(all_results)} RUNS COMPLETE\")\nprint(f\"{'='*60}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-19T09:06:17.541977Z","iopub.execute_input":"2026-09-19T09:06:17.542244Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 4, Cell 7: Compile main results table\n\n# Collect all results (reload from disk if needed)\nresults_files = sorted(LOG_DIR.glob('*_results.json'))\nall_results_loaded = []\nfor rf in results_files:\n    with open(rf, 'r') as f:\n        all_results_loaded.append(json.load(f))\n\n# Build table\nrows = []\nfor r in all_results_loaded:\n    rows.append({\n        'run_name': r['run_name'],\n        'head_type': r['config']['head_type'],\n        'seed': r['config']['seed'],\n        'seen_score': r['seen']['score'],\n        'seen_root': r['seen']['root'],\n        'seen_vowel': r['seen']['vowel'],\n        'seen_cons': r['seen']['cons'],\n        'seen_whole_acc': r['seen']['whole_acc'],\n        'unseen_score': r['unseen']['score'],\n        'unseen_root': r['unseen']['root'],\n        'unseen_vowel': r['unseen']['vowel'],\n        'unseen_cons': r['unseen']['cons'],\n        'unseen_whole_acc': r['unseen']['whole_acc'],\n        'gap': r['gap'],\n        'best_epoch': r['best_epoch'],\n        'time_min': r['total_time_min'],\n    })\n\nresults_df = pd.DataFrame(rows)\nresults_df.to_csv(LOG_DIR / 'main_experiment_table.csv', index=False)\n\n# Display\nprint(\"\\n=== MAIN EXPERIMENT RESULTS ===\\n\")\ndisplay_cols = ['run_name', 'head_type', 'seed', 'seen_score', 'unseen_score', 'gap', 'best_epoch']\nprint(results_df[display_cols].to_string(index=False, float_format='%.4f'))\n\n# Aggregate by head type\nprint(\"\\n=== AGGREGATED (mean ± std across seeds) ===\\n\")\nfor ht in ['component', 'flat']:\n    sub = results_df[results_df['head_type'] == ht]\n    if len(sub) == 0:\n        continue\n    print(f\"  {ht.upper()}\")\n    print(f\"    Seen score:     {sub['seen_score'].mean():.4f} ± {sub['seen_score'].std():.4f}\")\n    print(f\"    Unseen score:   {sub['unseen_score'].mean():.4f} ± {sub['unseen_score'].std():.4f}\")\n    print(f\"    Gap:            {sub['gap'].mean():.4f} ± {sub['gap'].std():.4f}\")\n    print(f\"    Seen whole acc: {sub['seen_whole_acc'].mean():.4f} ± {sub['seen_whole_acc'].std():.4f}\")\n    print(f\"    Unseen whole:   {sub['unseen_whole_acc'].mean():.4f} ± {sub['unseen_whole_acc'].std():.4f}\")\n    print()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 4, Cell 8: The headline figure — flat vs component on Seen vs Unseen\n\nresults_df_loaded = pd.read_csv(LOG_DIR / 'main_experiment_table.csv')\n\ncomp = results_df_loaded[results_df_loaded['head_type'] == 'component']\nflat = results_df_loaded[results_df_loaded['head_type'] == 'flat']\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5.5))\n\n# Panel 1: Bar chart — Seen vs Unseen, Component vs Flat\nx = np.arange(2)\nwidth = 0.3\n\ncomp_means = [comp['seen_score'].mean(), comp['unseen_score'].mean()]\ncomp_stds = [comp['seen_score'].std(), comp['unseen_score'].std()]\nflat_means = [flat['seen_score'].mean(), flat['unseen_score'].mean()]\nflat_stds = [flat['seen_score'].std(), flat['unseen_score'].std()]\n\nbars1 = axes[0].bar(x - width/2, comp_means, width, yerr=comp_stds,\n                     color='#1F6F6B', label='Component (3 heads)', capsize=5, alpha=0.9)\nbars2 = axes[0].bar(x + width/2, flat_means, width, yerr=flat_stds,\n                     color='#C0392B', label='Flat (1,036-way)', capsize=5, alpha=0.9)\n\n# Add value labels on bars\nfor bars in [bars1, bars2]:\n    for bar in bars:\n        h = bar.get_height()\n        axes[0].text(bar.get_x() + bar.get_width() / 2, h + 0.01,\n                     f'{h:.3f}', ha='center', va='bottom', fontsize=11, fontweight='bold')\n\naxes[0].set_xticks(x)\naxes[0].set_xticklabels(['Seen test', 'Unseen test'], fontsize=12)\naxes[0].set_ylabel('Hierarchical macro-recall', fontsize=12)\naxes[0].set_ylim(0, 1.05)\naxes[0].set_title('The main result: seen vs. unseen', fontsize=14)\naxes[0].legend(fontsize=11, loc='lower left')\naxes[0].grid(axis='y', alpha=0.3)\n\n# Panel 2: Gap comparison\ngap_comp = comp['gap'].values\ngap_flat = flat['gap'].values\n\nbp = axes[1].boxplot([gap_comp, gap_flat],\n                      labels=['Component', 'Flat'],\n                      patch_artist=True,\n                      widths=0.4)\nbp['boxes'][0].set_facecolor('#DCEAE8')\nbp['boxes'][0].set_edgecolor('#1F6F6B')\nbp['boxes'][1].set_facecolor('#F6E4E0')\nbp['boxes'][1].set_edgecolor('#C0392B')\n\naxes[1].set_ylabel('Seen − Unseen gap', fontsize=12)\naxes[1].set_title('Generalization gap (lower = better)', fontsize=14)\naxes[1].grid(axis='y', alpha=0.3)\n\n# Annotate means\nfor i, (data, color) in enumerate([(gap_comp, '#1F6F6B'), (gap_flat, '#C0392B')]):\n    axes[1].text(i + 1, np.mean(data) + 0.01, f'{np.mean(data):.3f}',\n                 ha='center', fontsize=11, fontweight='bold', color=color)\n\nplt.suptitle('Compositional decomposition vs. flat classification\\n'\n             'EfficientNet-B0 + GeM, 20% hold-out, 3 seeds',\n             fontsize=14, y=1.04)\nplt.tight_layout()\nplt.savefig(FIG_DIR / 'main_result_flat_vs_component.png', dpi=140, bbox_inches='tight')\nplt.show()\n\nprint(\"✓ Headline figure saved\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}