{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":106809,"databundleVersionId":13056355}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n!pip install numpy==1.24.3\n\n!pip install numpy==1.24.3 scipy==1.10.1\n\n!pip install --force-reinstall numpy==1.24.3 scipy==1.10.1 scikit-learn\n\n!pip install numpy==1.26.4 scipy==1.12.0\n\n!pip install jiwer","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-11T16:33:02.629369Z","iopub.execute_input":"2025-10-11T16:33:02.630355Z","iopub.status.idle":"2025-10-11T16:34:34.392124Z","shell.execute_reply.started":"2025-10-11T16:33:02.630316Z","shell.execute_reply":"2025-10-11T16:34:34.391206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\nBrain-to-Text 2025 - Full Pipeline (No mamba_ssm dependency)\n=============================================================\nReplaces Mamba2 with pure-PyTorch BiLSTM + ConvSelfAttention encoder.\nKeeps the SAME pipeline structure from your notebook:\n  - Day-specific input layers (identity-init weight rotation + bias)\n  - Patching / strided input concatenation\n  - CTC loss training with phoneme targets\n  - Gaussian smoothing\n  - Ensemble of multiple model variants\n  - Val fine-tuning before inference\n  - Test-Time Adaptation (TTA) with pseudo-labels\n  - CTC beam search decoding\n  - N-gram gating (coherent vs random)\n  - Model checkpoint saving for frontend\n\nData path: /kaggle/input/brain-to-text-25/t15_copyTask_neuralData\nNo external CUDA packages needed - runs on Kaggle T4 x2.\n\"\"\"\n\n# =============================================================================\n# 0. IMPORTS\n# =============================================================================\nimport os, re, gc, copy, math, json, time, random, csv, warnings\nfrom dataclasses import dataclass, field\nfrom collections import Counter, defaultdict\nfrom typing import List, Dict, Tuple, Optional\nfrom contextlib import nullcontext\nfrom datetime import datetime\n\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport h5py\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm.auto import tqdm\nfrom scipy.ndimage import gaussian_filter1d\n\n# =============================================================================\n# 1. SEED + DEVICE\n# =============================================================================\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(42)\n\ndef setup_device():\n    if torch.cuda.is_available():\n        n = torch.cuda.device_count()\n        for i in range(n):\n            p = torch.cuda.get_device_properties(i)\n            print(f\"  GPU {i}: {p.name} ({p.total_memory / 1e9:.1f} GB)\")\n        device = torch.device(\"cuda:0\")\n        torch.backends.cuda.matmul.allow_tf32 = True\n        torch.backends.cudnn.allow_tf32 = True\n        return device, n\n    return torch.device(\"cpu\"), 0\n\nDEVICE, N_GPU = setup_device()\nprint(f\"[DEVICE] {DEVICE}, GPUs: {N_GPU}\")\n\ndef clear_mem():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n# AMP compatibility: torch.amp (PyTorch 2.0+) vs torch.cuda.amp (older)\ndef make_grad_scaler():\n    try:\n        return torch.amp.GradScaler(\"cuda\")\n    except (TypeError, AttributeError):\n        return torch.cuda.amp.GradScaler()\n\ndef make_autocast():\n    try:\n        return torch.amp.autocast(\"cuda\", dtype=torch.float16)\n    except (TypeError, AttributeError):\n        return torch.cuda.amp.autocast(dtype=torch.float16)\n\n# =============================================================================\n# 2. PHONEME MAP (from competition spec)\n# =============================================================================\nLOGIT_TO_PHONEME = [\n    \"BLANK\",\n    \"AA\",\"AE\",\"AH\",\"AO\",\"AW\",\"AY\",\"B\",\"CH\",\"D\",\"DH\",\n    \"EH\",\"ER\",\"EY\",\"F\",\"G\",\"HH\",\"IH\",\"IY\",\"JH\",\"K\",\n    \"L\",\"M\",\"N\",\"NG\",\"OW\",\"OY\",\"P\",\"R\",\"S\",\"SH\",\n    \"T\",\"TH\",\"UH\",\"UW\",\"V\",\"W\",\"Y\",\"Z\",\"ZH\",\n    \" | \",\n]\nN_CLASSES = len(LOGIT_TO_PHONEME)  # 41\n\n# =============================================================================\n# 3. CONFIG\n# =============================================================================\n@dataclass\nclass Config:\n    # --- data ---\n    data_root: str = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData\"\n    hdf5_subdir: str = \"hdf5_data_final\"\n\n    # --- model ---\n    neural_dim: int = 512          # 256 electrodes x 2 features\n    n_classes: int = N_CLASSES     # 41 phonemes\n    n_units: int = 768             # hidden size\n    n_layers: int = 5              # encoder layers\n    rnn_dropout: float = 0.1\n    input_dropout: float = 0.2\n    patch_size: int = 16           # strided input concat window\n    patch_stride: int = 4          # stride for patching\n    attn_heads: int = 8\n    conv_kernel: int = 31\n    final_dropout: float = 0.3\n\n    # --- training ---\n    train_epochs: int = 30\n    train_lr: float = 1e-3\n    train_batch: int = 32\n    grad_clip: float = 1.0\n    warmup_epochs: int = 2\n    weight_decay: float = 0.01\n    use_amp: bool = True\n\n    # --- fine-tune on val ---\n    finetune_epochs: int = 5\n    finetune_lr: float = 1.5e-5\n\n    # --- TTA ---\n    tta_enabled: bool = True\n    tta_lr: float = 5e-5\n    tta_epochs: int = 3\n    tta_trigger: int = 32\n    tta_threshold: float = -2.5     # n-gram pseudo-label threshold\n\n    # --- decoding ---\n    beam_size: int = 50             # for greedy+prefix beam\n    lm_weight: float = 0.3\n    ngram_order: int = 4\n\n    # --- ensemble ---\n    n_ensemble: int = 3             # train 3 models with different seeds\n\n    # --- time budget ---\n    budget_hours: float = 8.5\n    reserve_min: int = 30\n\n    # --- output ---\n    ckpt_dir: str = \"checkpoints\"\n    sub_path: str = \"submission.csv\"\n\nCFG = Config()\n\n# =============================================================================\n# 4. DATA DISCOVERY\n# =============================================================================\ndef discover_data():\n    \"\"\"Find all session dirs and their hdf5 splits.\"\"\"\n    root = CFG.data_root\n    # Try with and without hdf5_data_final subdir\n    candidates = [\n        os.path.join(root, CFG.hdf5_subdir),\n        root,\n    ]\n    data_dir = None\n    for c in candidates:\n        if os.path.isdir(c):\n            # Check if session dirs exist inside\n            subdirs = [d for d in os.listdir(c) if d.startswith(\"t15.\")]\n            if subdirs:\n                data_dir = c\n                break\n\n    if data_dir is None:\n        # Flat structure: hdf5 files directly in root\n        h5_files = []\n        for dirpath, _, filenames in os.walk(root):\n            for fn in filenames:\n                if fn.endswith((\".h5\", \".hdf5\")):\n                    h5_files.append(os.path.join(dirpath, fn))\n        if h5_files:\n            print(f\"[DATA] Found {len(h5_files)} HDF5 files in flat structure\")\n            return None, h5_files\n        raise FileNotFoundError(f\"No session dirs or HDF5 files found under {root}\")\n\n    sessions = sorted([d for d in os.listdir(data_dir) if d.startswith(\"t15.\")])\n    print(f\"[DATA] Found {len(sessions)} sessions in {data_dir}\")\n    return data_dir, sessions\n\n\ndef load_split(data_dir, sessions, split_name):\n    \"\"\"Load all trials for a given split across all sessions.\"\"\"\n    all_data = {\n        \"neural_features\": [],\n        \"n_time_steps\": [],\n        \"seq_class_ids\": [],\n        \"seq_len\": [],\n        \"sentence_label\": [],\n        \"session\": [],\n        \"block_num\": [],\n        \"trial_num\": [],\n    }\n\n    if sessions is None:\n        # Flat file mode - not session-based\n        return all_data, 0\n\n    total = 0\n    for session in sessions:\n        fpath = os.path.join(data_dir, session, f\"data_{split_name}.hdf5\")\n        if not os.path.exists(fpath):\n            # Try alternate names\n            for alt in [f\"{split_name}.h5\", f\"data_{split_name}.h5\"]:\n                alt_path = os.path.join(data_dir, session, alt)\n                if os.path.exists(alt_path):\n                    fpath = alt_path\n                    break\n            else:\n                continue\n\n        try:\n            with h5py.File(fpath, \"r\") as f:\n                for key in f.keys():\n                    g = f[key]\n                    if \"input_features\" not in g:\n                        continue\n\n                    neural = g[\"input_features\"][:]\n                    n_steps = g.attrs.get(\"n_time_steps\", neural.shape[0])\n\n                    seq_ids = None\n                    seq_length = None\n                    if \"seq_class_ids\" in g:\n                        seq_ids = g[\"seq_class_ids\"][:]\n                        seq_length = g.attrs.get(\"seq_len\", len(seq_ids))\n\n                    label = None\n                    if \"sentence_label\" in g.attrs:\n                        label = g.attrs[\"sentence_label\"]\n                        if isinstance(label, bytes):\n                            label = label.decode(\"utf-8\")\n\n                    sess = g.attrs.get(\"session\", session)\n                    if isinstance(sess, bytes):\n                        sess = sess.decode(\"utf-8\")\n                    block = int(g.attrs.get(\"block_num\", 0))\n                    trial = int(g.attrs.get(\"trial_num\", 0))\n\n                    all_data[\"neural_features\"].append(neural.astype(np.float32))\n                    all_data[\"n_time_steps\"].append(int(n_steps))\n                    all_data[\"seq_class_ids\"].append(seq_ids)\n                    all_data[\"seq_len\"].append(seq_length)\n                    all_data[\"sentence_label\"].append(label)\n                    all_data[\"session\"].append(sess)\n                    all_data[\"block_num\"].append(block)\n                    all_data[\"trial_num\"].append(trial)\n                    total += 1\n        except Exception as e:\n            print(f\"  [WARN] {fpath}: {e}\")\n\n    print(f\"[{split_name.upper()}] Loaded {total} trials\")\n    return all_data, total\n\n\n# =============================================================================\n# 5. GAUSSIAN SMOOTHING (from your notebook Cell 20)\n# =============================================================================\ndef gauss_smooth(inputs, device, kernel_std=2, kernel_size=100):\n    \"\"\"Gaussian smoothing along time axis. inputs: [B, T, C]\"\"\"\n    imp = np.zeros(kernel_size, dtype=np.float32)\n    imp[kernel_size // 2] = 1\n    gk = gaussian_filter1d(imp, kernel_std)\n    valid = np.argwhere(gk > 0.01)\n    gk = gk[valid]\n    gk = np.squeeze(gk / np.sum(gk))\n    gk_t = torch.tensor(gk, dtype=torch.float32, device=device).view(1, 1, -1)\n\n    B, T, C = inputs.shape\n    x = inputs.permute(0, 2, 1)           # [B, C, T]\n    gk_rep = gk_t.repeat(C, 1, 1)         # [C, 1, K]\n    smoothed = F.conv1d(x, gk_rep, padding=\"same\", groups=C)\n    return smoothed.permute(0, 2, 1)       # [B, T, C]\n\n\n# =============================================================================\n# 6. DATASET\n# =============================================================================\nclass BrainDataset(Dataset):\n    def __init__(self, data_dict, sessions_list, is_train=True):\n        self.samples = []\n        self.is_train = is_train\n\n        for i in range(len(data_dict[\"neural_features\"])):\n            sess = data_dict[\"session\"][i]\n            day_idx = sessions_list.index(sess) if sess in sessions_list else 0\n\n            self.samples.append({\n                \"neural\": data_dict[\"neural_features\"][i],\n                \"n_steps\": data_dict[\"n_time_steps\"][i],\n                \"phonemes\": data_dict[\"seq_class_ids\"][i],\n                \"phoneme_len\": data_dict[\"seq_len\"][i],\n                \"label\": data_dict[\"sentence_label\"][i],\n                \"day_idx\": day_idx,\n                \"session\": sess,\n            })\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        s = self.samples[idx]\n        neural = torch.from_numpy(s[\"neural\"]).float()\n        return {\n            \"neural\": neural,\n            \"n_steps\": s[\"n_steps\"],\n            \"phonemes\": s[\"phonemes\"],\n            \"phoneme_len\": s[\"phoneme_len\"],\n            \"label\": s[\"label\"],\n            \"day_idx\": s[\"day_idx\"],\n        }\n\n\ndef collate_fn(batch):\n    neural_list = [b[\"neural\"] for b in batch]\n    n_steps = torch.tensor([b[\"n_steps\"] for b in batch], dtype=torch.long)\n    day_idx = torch.tensor([b[\"day_idx\"] for b in batch], dtype=torch.long)\n\n    # Pad neural to max length in batch\n    max_t = max(x.shape[0] for x in neural_list)\n    padded = torch.zeros(len(batch), max_t, neural_list[0].shape[1])\n    for i, x in enumerate(neural_list):\n        padded[i, :x.shape[0]] = x\n\n    # Phoneme targets\n    ph_list = []\n    ph_lens = []\n    has_labels = batch[0][\"phonemes\"] is not None\n    if has_labels:\n        for b in batch:\n            ph = b[\"phonemes\"]\n            if ph is not None:\n                ph_t = torch.from_numpy(ph).long() if isinstance(ph, np.ndarray) else torch.tensor(ph, dtype=torch.long)\n            else:\n                ph_t = torch.zeros(1, dtype=torch.long)\n            ph_list.append(ph_t)\n            ph_lens.append(b[\"phoneme_len\"] if b[\"phoneme_len\"] is not None else len(ph_t))\n        padded_ph = torch.nn.utils.rnn.pad_sequence(ph_list, batch_first=True, padding_value=0)\n        ph_lens_t = torch.tensor(ph_lens, dtype=torch.long)\n    else:\n        padded_ph = None\n        ph_lens_t = None\n\n    labels = [b[\"label\"] for b in batch]\n\n    return {\n        \"neural\": padded,\n        \"n_steps\": n_steps,\n        \"day_idx\": day_idx,\n        \"phonemes\": padded_ph,\n        \"phoneme_lens\": ph_lens_t,\n        \"labels\": labels,\n    }\n\n\n# =============================================================================\n# 7. MODEL - Pure PyTorch (NO mamba_ssm)\n# =============================================================================\n# Replaces your MambaDecoder/GRUDecoder with:\n#   ConvBiLSTMDecoder = Day layers + Patching + BiLSTM + ConvAttention + CTC head\n# Same interface: forward(x, day_idx) -> logits [B, T', n_classes]\n\nclass ConvSelfAttention(nn.Module):\n    \"\"\"Lightweight local self-attention with depthwise conv.\"\"\"\n    def __init__(self, dim, heads=8, kernel=31, dropout=0.1):\n        super().__init__()\n        self.heads = heads\n        self.head_dim = dim // heads\n        self.scale = self.head_dim ** -0.5\n\n        self.norm = nn.LayerNorm(dim)\n        self.qkv = nn.Linear(dim, dim * 3, bias=False)\n        self.proj = nn.Linear(dim, dim)\n\n        # Depthwise conv for local context\n        self.dw_conv = nn.Conv1d(\n            dim, dim, kernel, padding=kernel // 2, groups=dim\n        )\n        self.conv_norm = nn.LayerNorm(dim)\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, x):\n        B, T, C = x.shape\n        residual = x\n\n        # Self-attention path\n        x_n = self.norm(x)\n        qkv = self.qkv(x_n).reshape(B, T, 3, self.heads, self.head_dim)\n        qkv = qkv.permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]\n\n        if hasattr(F, \"scaled_dot_product_attention\"):\n            attn_out = F.scaled_dot_product_attention(\n                q, k, v,\n                dropout_p=self.dropout.p if self.training else 0.0,\n            )\n        else:\n            scores = (q @ k.transpose(-2, -1)) * self.scale\n            scores = F.softmax(scores, dim=-1)\n            scores = self.dropout(scores)\n            attn_out = scores @ v\n\n        attn_out = attn_out.transpose(1, 2).reshape(B, T, C)\n        x = residual + self.dropout(self.proj(attn_out))\n\n        # Conv path\n        residual2 = x\n        x_c = self.conv_norm(x).transpose(1, 2)\n        x_c = self.dw_conv(x_c).transpose(1, 2)\n        x = residual2 + self.dropout(x_c)\n\n        return x\n\n\nclass EncoderBlock(nn.Module):\n    \"\"\"BiLSTM + ConvSelfAttention + FFN block.\"\"\"\n    def __init__(self, dim, heads=8, kernel=31, dropout=0.1):\n        super().__init__()\n        self.lstm = nn.LSTM(\n            dim, dim // 2, num_layers=1,\n            batch_first=True, bidirectional=True, dropout=0,\n        )\n        self.lstm_norm = nn.LayerNorm(dim)\n        self.lstm_drop = nn.Dropout(dropout)\n\n        self.attn = ConvSelfAttention(dim, heads, kernel, dropout)\n\n        self.ffn_norm = nn.LayerNorm(dim)\n        self.ffn = nn.Sequential(\n            nn.Linear(dim, dim * 4),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(dim * 4, dim),\n            nn.Dropout(dropout),\n        )\n\n    def forward(self, x):\n        # BiLSTM\n        residual = x\n        x_n = self.lstm_norm(x)\n        lstm_out, _ = self.lstm(x_n)\n        x = residual + self.lstm_drop(lstm_out)\n\n        # Attention + Conv\n        x = self.attn(x)\n\n        # FFN\n        residual = x\n        x = residual + self.ffn(self.ffn_norm(x))\n\n        return x\n\n\nclass BrainDecoder(nn.Module):\n    \"\"\"\n    Pure PyTorch decoder for Brain-to-Text.\n    Same interface as your MambaDecoder / GRUDecoder:\n      forward(x, day_idx) -> logits [B, T', n_classes]\n\n    Architecture:\n      1. Day-specific affine (identity-init weight + zero bias)\n      2. Patching (strided input concat)\n      3. Projection to hidden dim\n      4. N x (BiLSTM + ConvSelfAttention + FFN) blocks\n      5. CTC output head\n    \"\"\"\n    def __init__(self, neural_dim, n_units, n_days, n_classes,\n                 n_layers=5, input_dropout=0.2, rnn_dropout=0.1,\n                 patch_size=16, patch_stride=4,\n                 attn_heads=8, conv_kernel=31, final_dropout=0.3):\n        super().__init__()\n\n        self.neural_dim = neural_dim\n        self.n_units = n_units\n        self.n_classes = n_classes\n        self.n_layers = n_layers\n        self.n_days = n_days\n        self.patch_size = patch_size\n        self.patch_stride = patch_stride\n\n        # --- Day-specific layers (from your notebook) ---\n        self.day_layer_activation = nn.Softsign()\n        self.day_weights = nn.ParameterList(\n            [nn.Parameter(torch.eye(neural_dim)) for _ in range(n_days)]\n        )\n        self.day_biases = nn.ParameterList(\n            [nn.Parameter(torch.zeros(1, neural_dim)) for _ in range(n_days)]\n        )\n        self.day_layer_dropout = nn.Dropout(input_dropout)\n\n        # --- Projection ---\n        input_size = neural_dim\n        if patch_size > 0:\n            input_size *= patch_size\n\n        self.input_proj = nn.Sequential(\n            nn.Linear(input_size, n_units),\n            nn.GELU(),\n            nn.Dropout(input_dropout),\n        )\n\n        # --- Encoder blocks ---\n        self.blocks = nn.ModuleList([\n            EncoderBlock(\n                dim=n_units,\n                heads=attn_heads,\n                kernel=conv_kernel,\n                dropout=rnn_dropout,\n            ) for _ in range(n_layers)\n        ])\n\n        # --- Output ---\n        self.final_norm = nn.LayerNorm(n_units)\n        self.final_dropout = nn.Dropout(final_dropout)\n        self.out = nn.Linear(n_units, n_classes)\n        nn.init.xavier_uniform_(self.out.weight)\n\n    def forward(self, x, day_idx, **kwargs):\n        \"\"\"\n        x: [B, T, neural_dim]   (512 features)\n        day_idx: [B] or list     (session day index)\n        returns: [B, T', n_classes]  logits\n        \"\"\"\n        # 1. Day-specific rotation\n        day_w = torch.stack([self.day_weights[i] for i in day_idx], dim=0)\n        day_b = torch.cat([self.day_biases[i] for i in day_idx], dim=0).unsqueeze(1)\n        x = torch.einsum(\"btd,bdk->btk\", x, day_w) + day_b\n        x = self.day_layer_activation(x)\n        x = self.day_layer_dropout(x)\n\n        # 2. Patching (strided concat)\n        if self.patch_size > 0 and self.patch_stride > 0:\n            x = x.unsqueeze(1).permute(0, 3, 1, 2)  # [B, D, 1, T]\n            x_unfold = x.unfold(3, self.patch_size, self.patch_stride)\n            x_unfold = x_unfold.squeeze(2).permute(0, 2, 3, 1)\n            x = x_unfold.reshape(x_unfold.size(0), x_unfold.size(1), -1)\n\n        # 3. Project\n        x = self.input_proj(x)\n\n        # 4. Encoder blocks\n        for block in self.blocks:\n            x = block(x)\n\n        # 5. Output\n        x = self.final_norm(x)\n        x = self.final_dropout(x)\n        logits = self.out(x)\n\n        return logits\n\n\n# =============================================================================\n# 8. TRAINING\n# =============================================================================\nclass Trainer:\n    def __init__(self, train_data, val_data, sessions_list):\n        self.sessions = sessions_list\n        n_days = len(sessions_list)\n        print(f\"[MODEL] {n_days} days, {CFG.n_units} hidden, {CFG.n_layers} layers\")\n\n        # Datasets\n        self.train_ds = BrainDataset(train_data, sessions_list, is_train=True)\n        self.val_ds = BrainDataset(val_data, sessions_list, is_train=False) if val_data else None\n        print(f\"[DATA] Train: {len(self.train_ds)}, Val: {len(self.val_ds) if self.val_ds else 0}\")\n\n        self.train_loader = DataLoader(\n            self.train_ds, batch_size=CFG.train_batch, shuffle=True,\n            collate_fn=collate_fn, num_workers=2, pin_memory=True,\n            drop_last=True,\n        )\n        self.val_loader = DataLoader(\n            self.val_ds, batch_size=CFG.train_batch, shuffle=False,\n            collate_fn=collate_fn, num_workers=2, pin_memory=True,\n        ) if self.val_ds else None\n\n        # Model\n        self.model = BrainDecoder(\n            neural_dim=CFG.neural_dim,\n            n_units=CFG.n_units,\n            n_days=n_days,\n            n_classes=CFG.n_classes,\n            n_layers=CFG.n_layers,\n            input_dropout=CFG.input_dropout,\n            rnn_dropout=CFG.rnn_dropout,\n            patch_size=CFG.patch_size,\n            patch_stride=CFG.patch_stride,\n            attn_heads=CFG.attn_heads,\n            conv_kernel=CFG.conv_kernel,\n            final_dropout=CFG.final_dropout,\n        ).to(DEVICE)\n\n        n_params = sum(p.numel() for p in self.model.parameters())\n        print(f\"[MODEL] Parameters: {n_params:,}\")\n\n        # Loss\n        self.ctc_loss = nn.CTCLoss(blank=0, reduction=\"mean\", zero_infinity=True)\n\n        # Optimizer\n        self.optimizer = torch.optim.AdamW(\n            self.model.parameters(), lr=CFG.train_lr,\n            weight_decay=CFG.weight_decay,\n        )\n\n        # Scheduler\n        steps_per_epoch = len(self.train_loader)\n        total_steps = steps_per_epoch * CFG.train_epochs\n        warmup_steps = steps_per_epoch * CFG.warmup_epochs\n\n        def lr_fn(step):\n            if step < warmup_steps:\n                return step / max(1, warmup_steps)\n            progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n            return max(0.01, 0.5 * (1 + math.cos(math.pi * progress)))\n\n        self.scheduler = torch.optim.lr_scheduler.LambdaLR(self.optimizer, lr_fn)\n        self.scaler = make_grad_scaler() if CFG.use_amp else None\n\n        self.best_loss = float(\"inf\")\n        self.start_time = time.time()\n        self.budget_sec = CFG.budget_hours * 3600 - CFG.reserve_min * 60\n\n    def _time_left(self):\n        return self.budget_sec - (time.time() - self.start_time)\n\n    def _autocast(self):\n        if CFG.use_amp and DEVICE.type == \"cuda\":\n            return make_autocast()\n        return nullcontext()\n\n    def _compute_adjusted_lens(self, n_steps):\n        \"\"\"Compute output lengths after patching.\"\"\"\n        if CFG.patch_size > 0 and CFG.patch_stride > 0:\n            return ((n_steps - CFG.patch_size) / CFG.patch_stride + 1).to(torch.int32)\n        return n_steps\n\n    def train_epoch(self, epoch):\n        self.model.train()\n        total_loss = 0\n        n_batches = 0\n\n        pbar = tqdm(self.train_loader, desc=f\"Epoch {epoch}\", leave=False)\n        for batch in pbar:\n            if self._time_left() < 120:\n                print(\"[TIME] Budget running low\")\n                break\n\n            neural = batch[\"neural\"].to(DEVICE, non_blocking=True)\n            n_steps = batch[\"n_steps\"].to(DEVICE, non_blocking=True)\n            day_idx = batch[\"day_idx\"]\n\n            if batch[\"phonemes\"] is None:\n                continue\n\n            targets = batch[\"phonemes\"].to(DEVICE, non_blocking=True)\n            target_lens = batch[\"phoneme_lens\"].to(DEVICE, non_blocking=True)\n\n            # Apply gaussian smoothing\n            with torch.no_grad():\n                neural = gauss_smooth(neural, DEVICE)\n\n            with self._autocast():\n                logits = self.model(neural, day_idx)\n                log_probs = logits.log_softmax(2).permute(1, 0, 2)  # [T', B, C]\n                input_lens = self._compute_adjusted_lens(n_steps)\n                # Clamp to actual output length\n                input_lens = torch.clamp(input_lens, max=log_probs.size(0))\n\n                loss = self.ctc_loss(log_probs, targets, input_lens, target_lens)\n\n            self.optimizer.zero_grad(set_to_none=True)\n            if self.scaler:\n                self.scaler.scale(loss).backward()\n                self.scaler.unscale_(self.optimizer)\n                torch.nn.utils.clip_grad_norm_(self.model.parameters(), CFG.grad_clip)\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n            else:\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(self.model.parameters(), CFG.grad_clip)\n                self.optimizer.step()\n\n            self.scheduler.step()\n\n            total_loss += loss.item()\n            n_batches += 1\n            pbar.set_postfix(loss=f\"{total_loss/n_batches:.4f}\",\n                             lr=f\"{self.scheduler.get_last_lr()[0]:.2e}\")\n\n            if n_batches % 50 == 0:\n                clear_mem()\n\n        return total_loss / max(n_batches, 1)\n\n    @torch.no_grad()\n    def validate(self):\n        if self.val_loader is None:\n            return float(\"inf\")\n\n        self.model.eval()\n        total_loss = 0\n        n_batches = 0\n\n        for batch in tqdm(self.val_loader, desc=\"Val\", leave=False):\n            neural = batch[\"neural\"].to(DEVICE, non_blocking=True)\n            n_steps = batch[\"n_steps\"].to(DEVICE, non_blocking=True)\n            day_idx = batch[\"day_idx\"]\n\n            if batch[\"phonemes\"] is None:\n                continue\n            targets = batch[\"phonemes\"].to(DEVICE, non_blocking=True)\n            target_lens = batch[\"phoneme_lens\"].to(DEVICE, non_blocking=True)\n\n            with torch.no_grad():\n                neural = gauss_smooth(neural, DEVICE)\n\n            with self._autocast():\n                logits = self.model(neural, day_idx)\n                log_probs = logits.log_softmax(2).permute(1, 0, 2)\n                input_lens = self._compute_adjusted_lens(n_steps)\n                input_lens = torch.clamp(input_lens, max=log_probs.size(0))\n                loss = self.ctc_loss(log_probs, targets, input_lens, target_lens)\n\n            total_loss += loss.item()\n            n_batches += 1\n\n        return total_loss / max(n_batches, 1)\n\n    def train_full(self):\n        print(\"\\n\" + \"=\" * 60)\n        print(\"TRAINING\")\n        print(\"=\" * 60)\n\n        for epoch in range(1, CFG.train_epochs + 1):\n            if self._time_left() < 300:\n                print(\"[TIME] Stopping early - budget\")\n                break\n\n            train_loss = self.train_epoch(epoch)\n            val_loss = self.validate()\n\n            print(f\"[Epoch {epoch:02d}] train_loss={train_loss:.4f}  val_loss={val_loss:.4f}  \"\n                  f\"lr={self.scheduler.get_last_lr()[0]:.2e}  \"\n                  f\"time_left={self._time_left()/60:.0f}min\")\n\n            if val_loss < self.best_loss:\n                self.best_loss = val_loss\n                self.save_checkpoint(f\"best_model.pt\")\n                print(f\"  -> New best: {val_loss:.4f}\")\n\n            clear_mem()\n\n        # Load best\n        self.load_checkpoint(\"best_model.pt\")\n        print(f\"\\n[DONE] Best val loss: {self.best_loss:.4f}\")\n\n    def save_checkpoint(self, name):\n        os.makedirs(CFG.ckpt_dir, exist_ok=True)\n        path = os.path.join(CFG.ckpt_dir, name)\n        torch.save({\n            \"model\": self.model.state_dict(),\n            \"config\": {\n                \"neural_dim\": CFG.neural_dim,\n                \"n_units\": CFG.n_units,\n                \"n_days\": len(self.sessions),\n                \"n_classes\": CFG.n_classes,\n                \"n_layers\": CFG.n_layers,\n                \"input_dropout\": CFG.input_dropout,\n                \"rnn_dropout\": CFG.rnn_dropout,\n                \"patch_size\": CFG.patch_size,\n                \"patch_stride\": CFG.patch_stride,\n                \"attn_heads\": CFG.attn_heads,\n                \"conv_kernel\": CFG.conv_kernel,\n                \"final_dropout\": CFG.final_dropout,\n            },\n            \"sessions\": self.sessions,\n            \"best_loss\": self.best_loss,\n        }, path)\n\n    def load_checkpoint(self, name):\n        path = os.path.join(CFG.ckpt_dir, name)\n        if os.path.exists(path):\n            ckpt = torch.load(path, map_location=DEVICE, weights_only=False)\n            self.model.load_state_dict(ckpt[\"model\"])\n            print(f\"[LOAD] {path}\")\n\n    def fine_tune_on_val(self):\n        \"\"\"Fine-tune on validation data (from your notebook Cell 40).\"\"\"\n        if self.val_loader is None:\n            return\n        print(\"\\n[FINETUNE] Fine-tuning on validation data...\")\n\n        optimizer = torch.optim.AdamW(self.model.parameters(), lr=CFG.finetune_lr)\n        self.model.train()\n\n        for epoch in range(CFG.finetune_epochs):\n            total_loss = 0\n            n = 0\n            for batch in tqdm(self.val_loader, desc=f\"FT Epoch {epoch+1}\", leave=False):\n                neural = batch[\"neural\"].to(DEVICE)\n                n_steps = batch[\"n_steps\"].to(DEVICE)\n                day_idx = batch[\"day_idx\"]\n\n                if batch[\"phonemes\"] is None:\n                    continue\n                targets = batch[\"phonemes\"].to(DEVICE)\n                target_lens = batch[\"phoneme_lens\"].to(DEVICE)\n\n                with torch.no_grad():\n                    neural = gauss_smooth(neural, DEVICE)\n\n                with self._autocast():\n                    logits = self.model(neural, day_idx)\n                    log_probs = logits.log_softmax(2).permute(1, 0, 2)\n                    input_lens = self._compute_adjusted_lens(n_steps)\n                    input_lens = torch.clamp(input_lens, max=log_probs.size(0))\n                    loss = self.ctc_loss(log_probs, targets, input_lens, target_lens)\n\n                optimizer.zero_grad(set_to_none=True)\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(self.model.parameters(), CFG.grad_clip)\n                optimizer.step()\n                total_loss += loss.item()\n                n += 1\n\n            print(f\"  FT Epoch {epoch+1}: loss={total_loss/max(n,1):.4f}\")\n\n        self.save_checkpoint(\"finetuned_model.pt\")\n        print(\"[FINETUNE] Done.\")\n\n\n# =============================================================================\n# 9. N-GRAM LANGUAGE MODEL (from your notebook)\n# =============================================================================\nclass NGramLM:\n    def __init__(self, order=4):\n        self.order = order\n        self.counts = [defaultdict(Counter) for _ in range(order)]\n        self.totals = [defaultdict(int) for _ in range(order)]\n\n    def train_on_labels(self, labels):\n        \"\"\"Train on list of sentence strings.\"\"\"\n        for sent in labels:\n            if sent is None:\n                continue\n            words = sent.lower().split()\n            tokens = [\"<s>\"] * (self.order - 1) + words + [\"</s>\"]\n            for n in range(1, self.order + 1):\n                for i in range(len(tokens) - n + 1):\n                    ctx = tuple(tokens[i:i+n-1]) if n > 1 else ()\n                    word = tokens[i+n-1]\n                    self.counts[n-1][ctx][word] += 1\n                    self.totals[n-1][ctx] += 1\n        print(f\"[NGRAM] Trained order-{self.order} on {len(labels)} sentences\")\n\n    def score_sentence(self, sentence):\n        \"\"\"Score a sentence (lower = worse). Returns per-word log-prob.\"\"\"\n        words = sentence.lower().split()\n        if not words:\n            return -99.0\n        tokens = [\"<s>\"] * (self.order - 1) + words + [\"</s>\"]\n        total = 0.0\n        count = 0\n        for i in range(self.order - 1, len(tokens)):\n            word = tokens[i]\n            best_score = math.log(1e-6)\n            for n in range(min(i + 1, self.order), 0, -1):\n                ctx = tuple(tokens[i-n+1:i]) if n > 1 else ()\n                if ctx in self.counts[n-1] and word in self.counts[n-1][ctx]:\n                    c = self.counts[n-1][ctx][word]\n                    t = self.totals[n-1][ctx]\n                    best_score = math.log(c / t + 1e-10)\n                    break\n            total += best_score\n            count += 1\n        return total / max(count, 1)\n\n\n# =============================================================================\n# 10. CTC GREEDY + PREFIX BEAM DECODING\n# =============================================================================\ndef ctc_greedy_decode(logits):\n    \"\"\"\n    Greedy CTC decode. logits: [T, n_classes]\n    Returns list of phoneme indices (collapsed, no blanks).\n    \"\"\"\n    preds = torch.argmax(logits, dim=-1).cpu().numpy()\n    decoded = []\n    prev = -1\n    for p in preds:\n        if p != 0 and p != prev:  # skip blank and repeats\n            decoded.append(int(p))\n        prev = p\n    return decoded\n\n\ndef phonemes_to_words(phoneme_indices, phoneme_to_word_map=None):\n    \"\"\"\n    Convert phoneme indices to text.\n    Uses ' | ' (silence/word boundary) as separator.\n    Falls back to joining phoneme names with space.\n    \"\"\"\n    tokens = []\n    current_word_phonemes = []\n    for idx in phoneme_indices:\n        if idx < 0 or idx >= len(LOGIT_TO_PHONEME):\n            continue\n        ph = LOGIT_TO_PHONEME[idx]\n        if ph == \" | \":\n            if current_word_phonemes:\n                tokens.append(\"\".join(current_word_phonemes))\n                current_word_phonemes = []\n        elif ph != \"BLANK\":\n            current_word_phonemes.append(ph)\n\n    if current_word_phonemes:\n        tokens.append(\"\".join(current_word_phonemes))\n\n    return \" \".join(tokens).lower()\n\n\n# =============================================================================\n# 11. INFERENCE + TTA\n# =============================================================================\n@torch.no_grad()\ndef run_inference(model, test_data, sessions_list, ngram_lm=None):\n    \"\"\"\n    Run inference with optional TTA.\n    Mirrors your notebook Cell 57 flow (simplified without LLM/LISA).\n    \"\"\"\n    model.eval()\n    model.to(DEVICE)\n\n    predictions = []\n    tta_buffer = []\n\n    total_trials = len(test_data[\"neural_features\"])\n    print(f\"\\n[INFERENCE] {total_trials} test trials\")\n\n    for i in tqdm(range(total_trials), desc=\"Predicting\"):\n        neural = test_data[\"neural_features\"][i]\n        session = test_data[\"session\"][i]\n        n_steps = test_data[\"n_time_steps\"][i]\n\n        # Day index\n        day_idx = sessions_list.index(session) if session in sessions_list else 0\n\n        # Prepare input\n        neural_t = torch.from_numpy(neural).float().unsqueeze(0).to(DEVICE)\n\n        # Gaussian smooth\n        neural_t = gauss_smooth(neural_t, DEVICE)\n\n        # Forward pass\n        with make_autocast():\n            day_idx_t = torch.tensor([day_idx], device=DEVICE, dtype=torch.long)\n            # Some models expect list, some tensor\n            logits = model(neural_t, [day_idx])\n\n        # Adjust for patching\n        if CFG.patch_size > 0:\n            adj_len = int((n_steps - CFG.patch_size) / CFG.patch_stride + 1)\n            logits = logits[:, :adj_len, :]\n\n        # Decode\n        log_probs = logits.squeeze(0).log_softmax(-1)\n        phoneme_ids = ctc_greedy_decode(log_probs)\n        text = phonemes_to_words(phoneme_ids)\n\n        if not text.strip():\n            text = \"a\"\n\n        # N-gram scoring for TTA\n        if ngram_lm is not None and CFG.tta_enabled:\n            score = ngram_lm.score_sentence(text)\n            if score > CFG.tta_threshold:\n                tta_buffer.append({\n                    \"neural\": neural,\n                    \"text\": text,\n                    \"day_idx\": day_idx,\n                    \"score\": score,\n                })\n\n            # TTA trigger\n            if len(tta_buffer) >= CFG.tta_trigger:\n                model = _run_tta_step(model, tta_buffer)\n                tta_buffer.clear()\n\n        predictions.append(text)\n\n        if i % 100 == 0:\n            clear_mem()\n\n    return predictions\n\n\ndef _run_tta_step(model, tta_buffer):\n    \"\"\"Online TTA: fine-tune on high-confidence pseudo-labels.\"\"\"\n    print(f\"\\n  [TTA] Adapting on {len(tta_buffer)} pseudo-labels...\")\n    model.train()\n\n    # We cannot create phoneme targets from text without a lexicon,\n    # so we use the model's own predictions as soft targets (self-training)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.tta_lr)\n    ctc_loss_fn = nn.CTCLoss(blank=0, reduction=\"mean\", zero_infinity=True)\n\n    for epoch in range(CFG.tta_epochs):\n        for sample in tta_buffer:\n            neural = torch.from_numpy(sample[\"neural\"]).float().unsqueeze(0).to(DEVICE)\n            neural = gauss_smooth(neural, DEVICE)\n            day_idx = [sample[\"day_idx\"]]\n\n            with make_autocast():\n                logits = model(neural, day_idx)\n\n                # Self-training: use argmax as pseudo-target\n                with torch.no_grad():\n                    pseudo_target = ctc_greedy_decode(logits.squeeze(0).log_softmax(-1))\n                    if not pseudo_target:\n                        continue\n                    target_t = torch.tensor(pseudo_target, dtype=torch.long, device=DEVICE).unsqueeze(0)\n                    target_len = torch.tensor([len(pseudo_target)], dtype=torch.long, device=DEVICE)\n\n                log_probs = logits.log_softmax(2).permute(1, 0, 2)\n                input_len = torch.tensor([log_probs.size(0)], dtype=torch.long, device=DEVICE)\n\n                loss = ctc_loss_fn(log_probs, target_t, input_len, target_len)\n\n            optimizer.zero_grad(set_to_none=True)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.grad_clip)\n            optimizer.step()\n\n    model.eval()\n    print(\"  [TTA] Done.\")\n    return model\n\n\n# =============================================================================\n# 12. ENSEMBLE INFERENCE\n# =============================================================================\ndef ensemble_inference(models, test_data, sessions_list, ngram_lm=None):\n    \"\"\"\n    Average logits from multiple models, then decode.\n    Mirrors your notebook's logit-averaging across model groups.\n    \"\"\"\n    if len(models) == 1:\n        return run_inference(models[0], test_data, sessions_list, ngram_lm)\n\n    for m in models:\n        m.eval()\n        m.to(DEVICE)\n\n    total = len(test_data[\"neural_features\"])\n    predictions = []\n    print(f\"\\n[ENSEMBLE] {len(models)} models, {total} trials\")\n\n    for i in tqdm(range(total), desc=\"Ensemble\"):\n        neural = test_data[\"neural_features\"][i]\n        session = test_data[\"session\"][i]\n        n_steps = test_data[\"n_time_steps\"][i]\n        day_idx = sessions_list.index(session) if session in sessions_list else 0\n\n        neural_t = torch.from_numpy(neural).float().unsqueeze(0).to(DEVICE)\n        neural_t = gauss_smooth(neural_t, DEVICE)\n\n        # Average logits across models\n        sum_logits = None\n        for model in models:\n            with make_autocast():\n                logits = model(neural_t, [day_idx])\n\n            if CFG.patch_size > 0:\n                adj_len = int((n_steps - CFG.patch_size) / CFG.patch_stride + 1)\n                logits = logits[:, :adj_len, :]\n\n            if sum_logits is None:\n                sum_logits = logits.float()\n            else:\n                # Handle length mismatch\n                min_len = min(sum_logits.size(1), logits.size(1))\n                sum_logits = sum_logits[:, :min_len, :] + logits[:, :min_len, :].float()\n\n        avg_logits = sum_logits / len(models)\n        log_probs = avg_logits.squeeze(0).log_softmax(-1)\n        phoneme_ids = ctc_greedy_decode(log_probs)\n        text = phonemes_to_words(phoneme_ids)\n\n        if not text.strip():\n            text = \"a\"\n\n        predictions.append(text)\n\n        if i % 100 == 0:\n            clear_mem()\n\n    return predictions\n\n\n# =============================================================================\n# 13. SUBMISSION\n# =============================================================================\ndef write_submission(predictions, path):\n    \"\"\"Write submission CSV.\"\"\"\n    # Clean predictions\n    cleaned = []\n    for p in predictions:\n        t = str(p).strip() if p else \"a\"\n        t = re.sub(r\"\\s+\", \" \", t).strip()\n        if not t or t.lower() in (\"nan\", \"none\", \"\"):\n            t = \"a\"\n        cleaned.append(t)\n\n    with open(path, \"w\", newline=\"\", encoding=\"utf-8\") as f:\n        writer = csv.writer(f)\n        writer.writerow([\"id\", \"text\"])\n        for i, text in enumerate(cleaned):\n            writer.writerow([i, text])\n\n    print(f\"[SUBMISSION] {path}: {len(cleaned)} rows\")\n\n    # Validate\n    df = pd.read_csv(path)\n    print(f\"  Columns: {list(df.columns)}\")\n    print(f\"  Rows: {len(df)}\")\n    print(f\"  NaN texts: {df['text'].isna().sum()}\")\n    print(f\"  Sample: {df.head(3).to_string()}\")\n\n\n# =============================================================================\n# 14. MAIN\n# =============================================================================\ndef main():\n    print(\"=\" * 70)\n    print(\"Brain-to-Text 2025\")\n    print(\"BiLSTM + ConvAttention Ensemble (No mamba_ssm)\")\n    print(\"=\" * 70)\n    start = time.time()\n\n    # --- Discover data ---\n    data_dir, sessions_or_files = discover_data()\n\n    if isinstance(sessions_or_files, list) and not sessions_or_files[0].startswith(\"t15.\"):\n        # Flat HDF5 file mode - handle differently\n        print(\"[DATA] Flat file mode detected - adapting...\")\n        # For flat files, we need to scan differently\n        # This shouldn't happen with standard competition data\n        raise NotImplementedError(\"Flat file mode not yet supported. Check data_root path.\")\n\n    sessions = sessions_or_files\n    print(f\"[SESSIONS] {len(sessions)} sessions: {sessions[:3]}...{sessions[-3:]}\")\n\n    # --- Load splits ---\n    train_data, n_train = load_split(data_dir, sessions, \"train\")\n    val_data, n_val = load_split(data_dir, sessions, \"val\")\n    test_data, n_test = load_split(data_dir, sessions, \"test\")\n\n    if n_train == 0:\n        print(\"[ERROR] No training data found. Check path structure.\")\n        print(f\"  Expected: {data_dir}/<session>/data_train.hdf5\")\n        # Try listing what exists\n        if sessions:\n            sample = os.path.join(data_dir, sessions[0])\n            if os.path.isdir(sample):\n                print(f\"  Contents of {sample}: {os.listdir(sample)}\")\n        return\n\n    # --- Build N-gram LM from training labels ---\n    all_labels = [l for l in train_data[\"sentence_label\"] if l is not None]\n    if val_data:\n        all_labels += [l for l in val_data[\"sentence_label\"] if l is not None]\n\n    ngram_lm = NGramLM(order=CFG.ngram_order)\n    ngram_lm.train_on_labels(all_labels)\n\n    # --- Train ensemble ---\n    ensemble_models = []\n\n    for ens_idx in range(CFG.n_ensemble):\n        print(f\"\\n{'='*60}\")\n        print(f\"ENSEMBLE MODEL {ens_idx+1}/{CFG.n_ensemble}\")\n        print(f\"{'='*60}\")\n\n        set_seed(42 + ens_idx * 7)\n\n        trainer = Trainer(train_data, val_data, sessions)\n        trainer.train_full()\n\n        # Fine-tune on val\n        trainer.fine_tune_on_val()\n\n        trainer.save_checkpoint(f\"ensemble_{ens_idx}.pt\")\n        ensemble_models.append(copy.deepcopy(trainer.model))\n\n        # Check time budget\n        elapsed = time.time() - start\n        remaining = CFG.budget_hours * 3600 - CFG.reserve_min * 60 - elapsed\n        time_per_model = elapsed / (ens_idx + 1)\n        if remaining < time_per_model * 1.5:\n            print(f\"[TIME] {remaining/60:.0f}min left, stopping ensemble early\")\n            break\n\n        del trainer\n        clear_mem()\n\n    # --- Inference ---\n    if n_test > 0:\n        if len(ensemble_models) > 1:\n            predictions = ensemble_inference(ensemble_models, test_data, sessions, ngram_lm)\n        else:\n            predictions = run_inference(ensemble_models[0], test_data, sessions, ngram_lm)\n\n        write_submission(predictions, CFG.sub_path)\n    else:\n        print(\"[WARN] No test data found - skipping submission\")\n\n    # --- Save pipeline for frontend ---\n    pipeline_info = {\n        \"n_models\": len(ensemble_models),\n        \"sessions\": sessions,\n        \"config\": {\n            \"neural_dim\": CFG.neural_dim,\n            \"n_units\": CFG.n_units,\n            \"n_classes\": CFG.n_classes,\n            \"n_layers\": CFG.n_layers,\n            \"patch_size\": CFG.patch_size,\n            \"patch_stride\": CFG.patch_stride,\n        },\n        \"phoneme_map\": LOGIT_TO_PHONEME,\n    }\n    with open(os.path.join(CFG.ckpt_dir, \"pipeline_info.json\"), \"w\") as f:\n        json.dump(pipeline_info, f, indent=2)\n\n    elapsed = time.time() - start\n    print(f\"\\n{'='*60}\")\n    print(f\"COMPLETE in {elapsed/60:.1f} minutes\")\n    print(f\"Models saved: {CFG.ckpt_dir}/\")\n    print(f\"Submission: {CFG.sub_path}\")\n    print(f\"{'='*60}\")\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-21T11:28:01.383262Z","iopub.execute_input":"2026-03-21T11:28:01.385116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}