{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":106809,"databundleVersionId":13056355}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"64690caa-0a0b-4049-970e-88a619b6d94c","cell_type":"markdown","source":"# DSD-NLA v2 · Part 1 — Model Architecture & Training\n\n> **Competition:** Brain-to-Text '25 (Kaggle)  \n> **Task:** Decode intracortical neural activity (512-ch, ~2 kHz) → natural language text  \n> **Metric:** Word Error Rate (**WER**, lower is better)  \n> **Dataset:** t15 copy-task neural recordings (HDF5, 512 channels)\n\n---\n\n## Abstract\n\nThis notebook implements **DSD-NLA v2**, an end-to-end brain-to-text decoder built on\nthe following 2024–2025 research insights:\n\n1. **Conformer Encoder** — convolution-augmented self-attention captures both local\n   temporal neural bursts and long-range dependencies (Gulati et al., 2020).\n2. **CTC Loss** as primary training objective — alignment-free, faster convergence\n   (Graves et al., 2006).\n3. **Hybrid CTC + Attention Decoder** — joint training: CTC guides encoder,\n   autoregressive decoder refines output.\n4. **Stochastic Depth + Label Smoothing + SpecAugment masking** — key regularisers\n   for small neural datasets (Brain-to-Text Benchmark '24 lessons).\n5. **OneCycleLR** with short warmup for stable, fast convergence.\n\n---\n\n## Architecture\n\n```\nNeural Signal  (B, T, 512)\n      │\n      ▼\nConv1d Subsampling  stride=2  →  (B, T//2, D)\n      │\n      ▼\nConformer Encoder  ×8 layers  D=512\n      │\n      ├──► CTC Head  →  CTC Loss\n      │\n      ▼\nAttention Decoder  ×6 layers (pre-LN Transformer)\n      │\n      ▼\nText (char-level, vocab=36)\n```\n\n---\n\n## References\n- Brain-to-Text Benchmark '24: Lessons Learned — https://arxiv.org/abs/2412.17227\n- BIT End-to-End Brain-to-Text — https://arxiv.org/abs/2511.21740\n- Conformer (Gulati et al., 2020)\n- CTC (Graves et al., 2006)","metadata":{}},{"id":"e423546a-ff7a-488e-83ec-f1c315537b05","cell_type":"markdown","source":"## Cell 1 — Imports & Environment","metadata":{}},{"id":"5d0d7220-5fa1-4a63-98a4-859a58366829","cell_type":"code","source":"\"\"\"\nCELL 1 — Imports & Environment\n\"\"\"\nimport os, re, math, string, glob\nfrom pathlib import Path\nfrom typing import List, Dict, Tuple, Optional, Any\nfrom dataclasses import dataclass\n\nimport numpy as np\nimport h5py\nimport pandas as pd\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence\n\nprint(f'PyTorch : {torch.__version__}')\nprint(f'CUDA    : {torch.cuda.is_available()}')\nif torch.cuda.is_available():\n    print(f'GPU     : {torch.cuda.get_device_name(0)}')\n    print(f'VRAM    : {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:32.928908Z","iopub.execute_input":"2026-03-07T11:45:32.929297Z","iopub.status.idle":"2026-03-07T11:45:32.935546Z","shell.execute_reply.started":"2026-03-07T11:45:32.929268Z","shell.execute_reply":"2026-03-07T11:45:32.934751Z"}},"outputs":[],"execution_count":null},{"id":"18c783cd-520c-41af-9acb-2b08e74177bf","cell_type":"markdown","source":"## Cell 2 — Configuration","metadata":{}},{"id":"94b1bd07-b3d1-43cf-9ddc-8fe4e9ccb1df","cell_type":"code","source":"\"\"\"\nCELL 2 — Configuration\n\nAll hyper-parameters live here. Change only this cell to experiment.\n\"\"\"\n@dataclass\nclass Config:\n    # ── Paths ──────────────────────────────────────────────────────\n    data_dir        : str   = '/kaggle/input/competitions/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final'\n    checkpoint_path : str   = '/kaggle/working/best_dsdnla_v2.pt'\n    submission_path : str   = '/kaggle/working/submission.csv'\n\n    # ── Data ───────────────────────────────────────────────────────\n    n_channels      : int   = 512      # electrode channels\n    max_neural_len  : int   = 1024     # clip long trials\n    max_text_len    : int   = 100      # max decode steps at inference\n\n    # ── Model ──────────────────────────────────────────────────────\n    d_model               : int   = 512\n    n_encoder_layers      : int   = 8\n    n_decoder_layers      : int   = 6\n    n_heads               : int   = 8\n    conv_kernel           : int   = 31    # Conformer depthwise conv kernel (must be odd)\n    ff_expansion          : int   = 4     # feedforward hidden = d_model * ff_expansion\n    dropout               : float = 0.15\n    stochastic_depth_prob : float = 0.1   # max drop-path probability\n    ctc_weight            : float = 0.3   # total = ctc_weight*CTC + (1-ctc_weight)*CE\n    label_smoothing       : float = 0.1\n\n    # ── Training ───────────────────────────────────────────────────\n    seed          : int   = 42\n    batch_size    : int   = 16\n    num_epochs    : int   = 60\n    learning_rate : float = 4e-4\n    weight_decay  : float = 1e-2\n    gradient_clip : float = 1.0\n    warmup_frac   : float = 0.05     # fraction of total steps for LR warmup\n    patience      : int   = 8        # early stopping patience\n    use_amp       : bool  = True     # automatic mixed precision\n\n    # ── Inference ──────────────────────────────────────────────────\n    ctc_decode : bool = True          # True=CTC greedy, False=attention decoder\n\n    # ── Device (set at runtime) ─────────────────────────────────────\n    device : str = 'cpu'\n\n\nCFG = Config()\nCFG.device = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n# Reproducibility\ntorch.manual_seed(CFG.seed)\nnp.random.seed(CFG.seed)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(CFG.seed)\n\nprint(f'Device : {CFG.device}')\nprint(CFG)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:32.936990Z","iopub.execute_input":"2026-03-07T11:45:32.937280Z","iopub.status.idle":"2026-03-07T11:45:32.953729Z","shell.execute_reply.started":"2026-03-07T11:45:32.937253Z","shell.execute_reply":"2026-03-07T11:45:32.953097Z"}},"outputs":[],"execution_count":null},{"id":"a1d9ad53-d3fe-4187-8cb2-d3932eba3069","cell_type":"markdown","source":"## Cell 3 — Data Structure Explorer\n\nRun this cell **once** to understand what is inside the HDF5 files.\nIt prints the folder tree and inspects the first trial of the first file.\n\n### What to expect\n```\nhdf5_data_final/\n  t15.2023.08.13/\n    data_train.hdf5   ← groups: 'trial_0', 'trial_1', ...\n      trial_0/\n        input_features  shape=(T, 512)  float32   ← neural spike-band power\n        attrs['sentence_label'] = 'the cat sat on the mat'\n```","metadata":{}},{"id":"a870c672-8a48-49e4-9b02-1ae3a6e0ffd5","cell_type":"code","source":"\"\"\"\nCELL 3 — Data Structure Explorer\n\nWalks the data directory and inspects the first HDF5 file found.\nSafe to run even if no files exist yet — prints a helpful error message.\n\"\"\"\ndef explore_data(data_dir: str):\n    data_path = Path(data_dir)\n    if not data_path.exists():\n        print(f'[ERROR] data_dir does not exist: {data_dir}')\n        print('  → Check that the Brain-to-Text 25 competition dataset is mounted.')\n        print('    Expected path: /kaggle/input/brain-to-text-25/...')\n        return\n\n    # List session folders\n    sessions = sorted(data_path.glob('t15.*'))\n    print(f'Found {len(sessions)} session folder(s) under {data_dir}')\n    for s in sessions[:5]:\n        files = list(s.glob('*.hdf5'))\n        print(f'  {s.name}/ → {[f.name for f in files]}')\n    if len(sessions) > 5:\n        print(f'  ... ({len(sessions) - 5} more)')\n\n    # Inspect first train file\n    train_files = sorted(data_path.glob('t15.*/data_train.hdf5'))\n    if not train_files:\n        print('[WARN] No data_train.hdf5 found. Data might not be mounted.')\n        return\n\n    print(f'\\nInspecting: {train_files[0]}')\n    with h5py.File(train_files[0], 'r') as f:\n        keys = list(f.keys())\n        print(f'  Top-level trial keys (first 5): {keys[:5]}  ... ({len(keys)} total)')\n        first_key = keys[0]\n        trial = f[first_key]\n        print(f'  Trial \"{first_key}\" datasets: {list(trial.keys())}')\n        if 'input_features' in trial:\n            feat = trial['input_features']\n            print(f'    input_features shape : {feat.shape}  dtype={feat.dtype}')\n            print(f'    input_features stats : mean={feat[:].mean():.4f}  std={feat[:].std():.4f}')\n        print(f'  Trial attributes: {dict(trial.attrs)}')\n\n\nexplore_data(CFG.data_dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:32.954730Z","iopub.execute_input":"2026-03-07T11:45:32.954991Z","iopub.status.idle":"2026-03-07T11:45:33.807842Z","shell.execute_reply.started":"2026-03-07T11:45:32.954958Z","shell.execute_reply":"2026-03-07T11:45:33.807165Z"}},"outputs":[],"execution_count":null},{"id":"8e18cc16-15c3-4cdb-b172-7fc003a319ac","cell_type":"markdown","source":"## Cell 4 — Character Tokenizer\n\n### Vocabulary (36 tokens)\n| ID | Token | Notes |\n|---|---|---|\n| 0 | `<blank>` | CTC blank — never appears in decoded output |\n| 1 | `<PAD>` | Padding — ignored by loss |\n| 2 | `<BOS>` | Begin of sequence — fed as first decoder token |\n| 3 | `<EOS>` | End of sequence — decoder stops when predicted |\n| 4–29 | `a–z` | Lowercase letters |\n| 30 | `<space>` | Word separator |\n| 31–35 | `' . , ! ?` | Basic punctuation |","metadata":{}},{"id":"072e62ad-5738-4c3a-8709-ea9657916c9e","cell_type":"code","source":"\"\"\"\nCELL 4 — Character-level Tokenizer\n\"\"\"\nclass CharTokenizer:\n    BLANK_ID    = 0   # CTC blank\n    PAD_ID      = 1   # padding\n    BOS_ID      = 2   # begin of sequence\n    EOS_ID      = 3   # end of sequence\n    SPECIAL_IDS = {0, 1, 2, 3}\n\n    def __init__(self):\n        self.chars: List[str] = ['<blank>', '<PAD>', '<BOS>', '<EOS>']\n        self.chars += list(string.ascii_lowercase)    # ids 4..29\n        self.chars += [' ', \"'\", '.', ',', '!', '?']  # ids 30..35\n\n        self.char2id: Dict[str, int] = {c: i for i, c in enumerate(self.chars)}\n        self.id2char: Dict[int, str] = {i: c for i, c in enumerate(self.chars)}\n        self.vocab_size: int = len(self.chars)  # 36\n\n    def encode(self, text: str, add_bos_eos: bool = True) -> torch.Tensor:\n        \"\"\"text → token-ID tensor, optionally wrapped with BOS/EOS.\"\"\"\n        text = text.lower().strip()\n        ids = [self.BOS_ID] if add_bos_eos else []\n        for ch in text:\n            ids.append(self.char2id.get(ch, self.char2id[' ']))\n        if add_bos_eos:\n            ids.append(self.EOS_ID)\n        return torch.tensor(ids, dtype=torch.long)\n\n    def encode_ctc(self, text: str) -> torch.Tensor:\n        \"\"\"\n        Encode WITHOUT BOS/EOS — used as CTC target.\n        CTC loss does NOT use BOS/EOS special tokens.\n        \"\"\"\n        return self.encode(text, add_bos_eos=False)\n\n    def decode(self, ids, remove_special: bool = True) -> str:\n        \"\"\"token-IDs → string (stops at EOS, optionally removes special tokens).\"\"\"\n        if isinstance(ids, torch.Tensor):\n            ids = ids.cpu().tolist()\n        chars = []\n        for i in ids:\n            if i == self.EOS_ID:\n                break\n            if remove_special and i in self.SPECIAL_IDS:\n                continue\n            ch = self.id2char.get(i, '')\n            if remove_special and ch.startswith('<'):\n                continue\n            chars.append(ch)\n        return ''.join(chars)\n\n    def ctc_decode_greedy(self, log_probs: torch.Tensor) -> List[str]:\n        \"\"\"\n        Greedy CTC decode.\n\n        Algorithm:\n          1. argmax over vocab at every time step → (B, T)\n          2. Collapse consecutive duplicate tokens\n          3. Remove BLANK tokens (id=0)\n          4. Decode remaining IDs to string\n\n        Args:\n            log_probs: (B, T, V)  — log-softmax output from CTC head\n        Returns:\n            List[str] of length B\n        \"\"\"\n        preds = log_probs.argmax(dim=-1)  # (B, T)\n        results = []\n        for seq in preds:\n            seq = seq.cpu().tolist()\n            # Step 2: collapse consecutive duplicates\n            collapsed, prev = [], -1\n            for t in seq:\n                if t != prev:\n                    collapsed.append(t)\n                prev = t\n            # Step 3: remove blanks\n            clean = [t for t in collapsed if t != self.BLANK_ID]\n            # Step 4: decode\n            results.append(self.decode(clean, remove_special=True))\n        return results\n\n\n# ── Tests ──────────────────────────────────────────────────────────\ntokenizer = CharTokenizer()\nprint(f'Vocab size : {tokenizer.vocab_size}')\n\nsample = 'hello world'\nenc    = tokenizer.encode(sample)\nprint(f\"encode('{sample}')     → {enc.tolist()}\")\nprint(f\"decode back            → '{tokenizer.decode(enc)}'\")\n\nctc_enc = tokenizer.encode_ctc(sample)\nprint(f\"encode_ctc('{sample}') → {ctc_enc.tolist()} (no BOS/EOS)\")\n\nassert tokenizer.decode(enc) == sample, 'Round-trip encode/decode failed'\nprint('Tokenizer tests passed.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:33.809327Z","iopub.execute_input":"2026-03-07T11:45:33.809757Z","iopub.status.idle":"2026-03-07T11:45:33.823103Z","shell.execute_reply.started":"2026-03-07T11:45:33.809730Z","shell.execute_reply":"2026-03-07T11:45:33.822466Z"}},"outputs":[],"execution_count":null},{"id":"4183c010-a9dc-4c8f-9df1-cec6b7127ba4","cell_type":"markdown","source":"## Cell 5 — Dataset & Augmentation\n\n### How the Dataset Works\n- Reads `input_features` (shape `T × 512`) from each trial group in the HDF5 files\n- Reads `sentence_label` attribute as the ground-truth text\n- Returns two targets per trial:\n  - `tokens` — BOS + char IDs + EOS (for the **attention decoder**)\n  - `ctc_tgt` — char IDs only (for **CTC loss**, no special tokens)\n\n### Augmentation (train only)\n| Technique | Probability | Effect |\n|---|---|---|\n| Time masking | 40% | Zero-out 1–2 windows of 5–20 timesteps |\n| Channel dropout | 25% | Drop 15% of electrode channels randomly |\n| Gaussian noise | 30% | Add σ=0.05 Gaussian noise |\n| Time stretch | 20% | Resample temporal axis by ±10% |","metadata":{}},{"id":"b0dd54b4-b8a9-486e-98fa-56dd0eaabdeb","cell_type":"code","source":"\"\"\"\nCELL 5 — BrainToTextDataset + collate_fn\n\"\"\"\nclass BrainToTextDataset(Dataset):\n    def __init__(\n        self,\n        hdf5_paths : List[str],\n        tokenizer  : CharTokenizer,\n        mode       : str  = 'train',\n        augment    : bool = True,\n        max_neural : int  = 1024,\n    ):\n        self.tok     = tokenizer\n        self.mode    = mode\n        self.augment = augment and (mode == 'train')\n        self.max_len = max_neural\n\n        self.trials: List[Tuple[int, str]] = []\n        self.files : Dict[int, h5py.File]  = {}\n\n        print(f'Loading {mode} data...')\n        for fi, path in enumerate(tqdm(hdf5_paths)):\n            f = h5py.File(path, 'r')\n            self.files[fi] = f\n            for key in f.keys():\n                self.trials.append((fi, key))\n        print(f'  → {len(self.trials)} trials loaded')\n\n    def __len__(self) -> int:\n        return len(self.trials)\n\n    def __getitem__(self, idx: int) -> Dict[str, Any]:\n        fi, key = self.trials[idx]\n        trial   = self.files[fi][key]\n\n        # ── Neural features (T, 512) ──────────────────────────────────\n        neural = torch.tensor(trial['input_features'][:], dtype=torch.float32)\n        if neural.size(0) > self.max_len:\n            neural = neural[: self.max_len]\n\n        # ── Text label ────────────────────────────────────────────────\n        text = trial.attrs.get('sentence_label', '') or ''\n        # h5py sometimes returns bytes or numpy bytes\n        if isinstance(text, (bytes, np.bytes_)):\n            text = text.decode('utf-8')\n        text = str(text).strip()\n\n        tokens  = self.tok.encode(text)      # [BOS, chars..., EOS]\n        ctc_tgt = self.tok.encode_ctc(text)  # [chars...] — no BOS/EOS for CTC\n\n        # Guard: CTC target must be non-empty (CTCLoss requires len >= 1)\n        if ctc_tgt.numel() == 0:\n            ctc_tgt = torch.tensor([self.tok.char2id[' ']], dtype=torch.long)\n\n        # ── Augmentation ──────────────────────────────────────────────\n        if self.augment:\n            neural = self._augment(neural)\n\n        return {\n            'neural'  : neural,    # (T, 512)\n            'tokens'  : tokens,    # (T_text,)  attention-decoder target\n            'ctc_tgt' : ctc_tgt,   # (T_ctc,)   CTC target\n            'text'    : text,\n        }\n\n    def _augment(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply data augmentation to neural features (train only).\"\"\"\n        T, C = x.shape\n\n        # 1) Time masking (SpecAugment-style)\n        if torch.rand(1).item() < 0.4:\n            n_masks = torch.randint(1, 3, (1,)).item()\n            for _ in range(int(n_masks)):\n                mask_len = torch.randint(5, 20, (1,)).item()\n                start    = torch.randint(0, max(1, T - mask_len), (1,)).item()\n                x[start : start + mask_len] = 0.0\n\n        # 2) Channel (electrode) dropout\n        if torch.rand(1).item() < 0.25:\n            keep = torch.bernoulli(torch.full((C,), 0.85))\n            x    = x * keep.unsqueeze(0)\n\n        # 3) Gaussian noise\n        if torch.rand(1).item() < 0.3:\n            x = x + torch.randn_like(x) * 0.05\n\n        # 4) Time sub-sampling (mild stretch/compress ±10%)\n        if torch.rand(1).item() < 0.2 and T > 10:\n            fac   = 0.9 + 0.2 * torch.rand(1).item()  # [0.9, 1.1]\n            new_T = max(10, int(T * fac))\n            idx   = torch.linspace(0, T - 1, new_T).long().clamp(0, T - 1)\n            x     = x[idx]\n\n        return x\n\n\ndef collate_fn(batch: List[Dict[str, Any]]) -> Dict[str, Any]:\n    \"\"\"\n    Collate a list of samples into a padded batch.\n\n    Returns a dict with:\n      neural       : (B, T_max, 512)  padded neural features\n      neural_mask  : (B, T_max) bool  True = PAD (encoder ignores these)\n      neural_lens  : (B,) int         original neural sequence lengths\n      tokens       : (B, T_text)      padded decoder targets (BOS+chars+EOS)\n      ctc_tgt      : (B, T_ctc)       padded CTC targets (chars only)\n      ctc_lens     : (B,) int         actual CTC target lengths\n      texts        : List[str]        raw text labels\n    \"\"\"\n    neurals  = [b['neural']   for b in batch]\n    tokens_l = [b['tokens']   for b in batch]\n    ctc_l    = [b['ctc_tgt']  for b in batch]\n    texts    = [b['text']     for b in batch]\n\n    neural_lens = torch.tensor([n.size(0) for n in neurals], dtype=torch.long)\n    ctc_lens    = torch.tensor([c.size(0) for c in ctc_l],   dtype=torch.long)\n\n    B     = len(neurals)\n    T_max = int(neural_lens.max().item())\n\n    # Pad neural sequences\n    neural_pad = pad_sequence(neurals, batch_first=True, padding_value=0.0)  # (B, T_max, 512)\n\n    # Build key-padding mask: True = PAD position\n    neural_mask = torch.ones(B, T_max, dtype=torch.bool)   # start: all PAD\n    for i in range(B):\n        neural_mask[i, : int(neural_lens[i].item())] = False  # valid → False\n\n    tokens_pad = pad_sequence(tokens_l, batch_first=True, padding_value=CharTokenizer.PAD_ID)\n    ctc_pad    = pad_sequence(ctc_l,    batch_first=True, padding_value=CharTokenizer.PAD_ID)\n\n    return {\n        'neural'      : neural_pad,\n        'neural_mask' : neural_mask,\n        'neural_lens' : neural_lens,\n        'tokens'      : tokens_pad,\n        'ctc_tgt'     : ctc_pad,\n        'ctc_lens'    : ctc_lens,\n        'texts'       : texts,\n    }\n\n\nprint('Dataset helpers defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:33.824119Z","iopub.execute_input":"2026-03-07T11:45:33.824329Z","iopub.status.idle":"2026-03-07T11:45:33.845101Z","shell.execute_reply.started":"2026-03-07T11:45:33.824310Z","shell.execute_reply":"2026-03-07T11:45:33.844410Z"}},"outputs":[],"execution_count":null},{"id":"2b6731e1-d989-469d-8dfb-f7558b754bfa","cell_type":"markdown","source":"## Cell 6 — Conformer Encoder\n\n### Why Conformer over plain Transformer?\nNeural spike-band power data has **both** local temporal structure\n(short bursts of activity) **and** long-range dependencies (sentence context).\nConformer adds a **depthwise convolution module** between the two feed-forward\nhalf-steps, capturing local structure better than attention alone.\n\n### Conformer Block structure\n```\nx → FF₁(½) → MHSA → ConvModule → FF₂(½) → LayerNorm → output\n     ↑ each sub-module has a residual connection\n```","metadata":{}},{"id":"6b4612ae-43a8-4a56-a09a-d1e36bb342d8","cell_type":"code","source":"\"\"\"\nCELL 6 — Conformer Encoder\n\"\"\"\nclass Swish(nn.Module):\n    \"\"\"Swish activation: x * sigmoid(x). Shown to work better than ReLU/GELU in Conformers.\"\"\"\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        return x * torch.sigmoid(x)\n\n\nclass PositionalEncoding(nn.Module):\n    \"\"\"Sinusoidal positional encoding (Vaswani et al., 2017).\"\"\"\n    def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 8000):\n        super().__init__()\n        self.dropout = nn.Dropout(dropout)\n        pe  = torch.zeros(max_len, d_model)\n        pos = torch.arange(0, max_len, dtype=torch.float32).unsqueeze(1)\n        div = torch.exp(\n            torch.arange(0, d_model, 2, dtype=torch.float32)\n            * (-math.log(10000.0) / d_model)\n        )\n        pe[:, 0::2] = torch.sin(pos * div)\n        pe[:, 1::2] = torch.cos(pos * div)\n        self.register_buffer('pe', pe.unsqueeze(0))  # (1, max_len, d_model)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        return self.dropout(x + self.pe[:, : x.size(1)])\n\n\nclass ConformerConvModule(nn.Module):\n    \"\"\"\n    Conformer depthwise-separable convolution module.\n\n    Structure: LayerNorm → Pointwise(2D) → GLU → Depthwise → BN → Swish → Pointwise → Dropout\n    The GLU gate halves the channels after the first pointwise expansion.\n    \"\"\"\n    def __init__(self, d_model: int, kernel_size: int = 31, dropout: float = 0.1):\n        super().__init__()\n        assert (kernel_size - 1) % 2 == 0, 'conv_kernel must be odd for symmetric padding'\n        self.ln   = nn.LayerNorm(d_model)\n        self.pw1  = nn.Conv1d(d_model, 2 * d_model, 1)     # pointwise expand ×2\n        self.glu  = nn.GLU(dim=1)                           # gate halves channels → d_model\n        self.dw   = nn.Conv1d(                              # depthwise conv (per-channel)\n            d_model, d_model, kernel_size,\n            padding=(kernel_size - 1) // 2,\n            groups=d_model,\n        )\n        self.bn   = nn.BatchNorm1d(d_model)\n        self.act  = Swish()\n        self.pw2  = nn.Conv1d(d_model, d_model, 1)         # pointwise project\n        self.drop = nn.Dropout(dropout)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        # x: (B, T, D)\n        residual = x\n        x = self.ln(x).transpose(1, 2)          # (B, D, T)\n        x = self.glu(self.pw1(x))               # (B, D, T)\n        x = self.act(self.bn(self.dw(x)))       # (B, D, T)\n        x = self.drop(self.pw2(x)).transpose(1, 2)  # (B, T, D)\n        return x + residual  # residual connection\n\n\nclass ConformerBlock(nn.Module):\n    \"\"\"\n    One Conformer block with stochastic depth.\n\n    BUG FIX from v2a: LayerNorm before MHSA is computed ONCE and reused\n    for Q, K, V.  Previously called self.mhsa_ln(x) three times independently\n    which is incorrect (different dropout masks would apply each time).\n    \"\"\"\n    def __init__(\n        self,\n        d_model      : int,\n        n_heads      : int,\n        ff_expansion : int   = 4,\n        conv_kernel  : int   = 31,\n        dropout      : float = 0.1,\n        drop_path    : float = 0.0,\n    ):\n        super().__init__()\n        d_ff = d_model * ff_expansion\n\n        self.ff1_ln  = nn.LayerNorm(d_model)\n        self.ff1     = nn.Sequential(\n            nn.Linear(d_model, d_ff), Swish(), nn.Dropout(dropout),\n            nn.Linear(d_ff, d_model), nn.Dropout(dropout),\n        )\n        self.mhsa_ln   = nn.LayerNorm(d_model)\n        self.mhsa      = nn.MultiheadAttention(d_model, n_heads, dropout=dropout, batch_first=True)\n        self.mhsa_drop = nn.Dropout(dropout)\n        self.conv      = ConformerConvModule(d_model, conv_kernel, dropout)\n        self.ff2_ln    = nn.LayerNorm(d_model)\n        self.ff2       = nn.Sequential(\n            nn.Linear(d_model, d_ff), Swish(), nn.Dropout(dropout),\n            nn.Linear(d_ff, d_model), nn.Dropout(dropout),\n        )\n        self.final_ln  = nn.LayerNorm(d_model)\n        self.drop_path = drop_path\n\n    def _drop(self, residual: torch.Tensor, shortcut: torch.Tensor) -> torch.Tensor:\n        \"\"\"Stochastic depth: randomly skip a residual branch during training.\"\"\"\n        if self.training and self.drop_path > 0.0:\n            survive = (torch.rand(1, device=residual.device) > self.drop_path).float()\n            return shortcut + residual * survive\n        return shortcut + residual\n\n    def forward(\n        self,\n        x                : torch.Tensor,\n        key_padding_mask : Optional[torch.Tensor] = None,\n    ) -> torch.Tensor:\n        # FF1 (half-step weight = 0.5)\n        x = self._drop(0.5 * self.ff1(self.ff1_ln(x)), x)\n\n        # MHSA — FIX: compute normed once, reuse for Q, K, V\n        normed = self.mhsa_ln(x)\n        attn_out, _ = self.mhsa(\n            normed, normed, normed,\n            key_padding_mask=key_padding_mask,\n            need_weights=False,\n        )\n        x = self._drop(self.mhsa_drop(attn_out), x)\n\n        # Conv module\n        x = self.conv(x)\n\n        # FF2 (half-step)\n        x = self._drop(0.5 * self.ff2(self.ff2_ln(x)), x)\n\n        return self.final_ln(x)\n\n\nclass ConformerEncoder(nn.Module):\n    \"\"\"\n    Full Conformer encoder.\n\n    Pipeline:\n      Conv subsampling (stride=2) → Positional Encoding → N × ConformerBlock\n\n    Returns:\n      z        : (B, T//2, D)  encoder output\n      enc_mask : (B, T//2)     padding mask (True = PAD)\n    \"\"\"\n    def __init__(\n        self,\n        n_channels   : int   = 512,\n        d_model      : int   = 512,\n        n_layers     : int   = 8,\n        n_heads      : int   = 8,\n        ff_expansion : int   = 4,\n        conv_kernel  : int   = 31,\n        dropout      : float = 0.1,\n        drop_path    : float = 0.1,\n    ):\n        super().__init__()\n        # Temporal subsampling: reduces sequence length by 2×\n        self.subsampler = nn.Sequential(\n            nn.Conv1d(n_channels, d_model, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(d_model),\n            nn.GELU(),\n            nn.Dropout(dropout),\n        )\n        self.pos_enc = PositionalEncoding(d_model, dropout, max_len=8000)\n\n        # Linearly increase stochastic depth from 0 → drop_path across layers\n        dpr = [drop_path * i / max(n_layers - 1, 1) for i in range(n_layers)]\n        self.blocks = nn.ModuleList([\n            ConformerBlock(d_model, n_heads, ff_expansion, conv_kernel, dropout, dp)\n            for dp in dpr\n        ])\n\n    def forward(\n        self,\n        x    : torch.Tensor,\n        mask : Optional[torch.Tensor] = None,\n    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:\n        # x: (B, T, C)  mask: (B, T) bool True=PAD\n        x = x.transpose(1, 2)   # (B, C, T)\n        x = self.subsampler(x)  # (B, D, T//2)\n        x = x.transpose(1, 2)   # (B, T//2, D)\n        L = x.size(1)\n\n        # Downsample mask to match subsampled length\n        if mask is not None:\n            mask_ds = mask[:, ::2]\n            if mask_ds.size(1) > L:\n                mask_ds = mask_ds[:, :L]\n            elif mask_ds.size(1) < L:\n                pad_cols = torch.ones(\n                    mask_ds.size(0), L - mask_ds.size(1),\n                    dtype=torch.bool, device=mask_ds.device,\n                )\n                mask_ds = torch.cat([mask_ds, pad_cols], dim=1)\n        else:\n            mask_ds = None\n\n        x = self.pos_enc(x)\n        for block in self.blocks:\n            x = block(x, key_padding_mask=mask_ds)\n\n        return x, mask_ds  # (B, L, D), (B, L)\n\n\nprint('ConformerEncoder defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:33.846109Z","iopub.execute_input":"2026-03-07T11:45:33.846714Z","iopub.status.idle":"2026-03-07T11:45:33.868046Z","shell.execute_reply.started":"2026-03-07T11:45:33.846691Z","shell.execute_reply":"2026-03-07T11:45:33.867024Z"}},"outputs":[],"execution_count":null},{"id":"125f2298-0921-46fc-90ed-b194f730a5ca","cell_type":"markdown","source":"## Cell 7 — Attention Decoder\n\n### Role\nDuring training the decoder uses **teacher forcing**: it sees the ground-truth\nprefix and predicts the next character.  At inference it runs **greedy autoregressive\ndecoding** (or CTC greedy decode from the encoder alone, which is faster).\n\n### Fix: causal mask dtype\nPreviously `generate_square_subsequent_mask` returned a `float` tensor with `-inf` values.\nPyTorch ≥ 2.0 warns about mixing `float` causal mask with `bool` padding mask.\nWe now use `torch.bool` throughout.","metadata":{}},{"id":"52597c6f-efb4-4604-bd13-c35bf76b9e55","cell_type":"code","source":"\"\"\"\nCELL 7 — Attention Decoder\n\"\"\"\ndef make_causal_mask(T: int, device: torch.device) -> torch.Tensor:\n    \"\"\"\n    Upper-triangular boolean causal mask.\n    Shape: (T, T)  — True = 'do NOT attend' (future tokens are masked).\n\n    Example for T=4:\n      [[F, T, T, T],\n       [F, F, T, T],\n       [F, F, F, T],\n       [F, F, F, F]]\n    \"\"\"\n    return torch.triu(torch.ones(T, T, dtype=torch.bool, device=device), diagonal=1)\n\n\nclass AttentionDecoder(nn.Module):\n    \"\"\"\n    Transformer decoder that cross-attends to encoder output z.\n\n    Uses pre-LN (norm_first=True) for training stability.\n    \"\"\"\n    def __init__(\n        self,\n        vocab_size   : int,\n        d_model      : int   = 512,\n        n_layers     : int   = 6,\n        n_heads      : int   = 8,\n        ff_expansion : int   = 4,\n        dropout      : float = 0.1,\n    ):\n        super().__init__()\n        self.embed   = nn.Embedding(vocab_size, d_model, padding_idx=CharTokenizer.PAD_ID)\n        self.pos_enc = PositionalEncoding(d_model, dropout, max_len=4000)\n\n        dec_layer = nn.TransformerDecoderLayer(\n            d_model         = d_model,\n            nhead           = n_heads,\n            dim_feedforward = d_model * ff_expansion,\n            dropout         = dropout,\n            activation      = 'gelu',\n            batch_first     = True,\n            norm_first      = True,  # pre-LN: more stable training\n        )\n        self.decoder = nn.TransformerDecoder(dec_layer, num_layers=n_layers)\n        self.proj    = nn.Linear(d_model, vocab_size)\n\n    def forward(\n        self,\n        tgt         : torch.Tensor,                    # (B, T_tgt) token ids\n        memory      : torch.Tensor,                    # (B, T_enc, D) encoder output\n        memory_mask : Optional[torch.Tensor] = None,   # (B, T_enc) True=PAD\n    ) -> torch.Tensor:                                 # → (B, T_tgt, V)\n        T = tgt.size(1)\n        causal  = make_causal_mask(T, tgt.device)      # (T, T) bool\n        tgt_pad = (tgt == CharTokenizer.PAD_ID)        # (B, T_tgt) bool\n\n        emb = self.pos_enc(self.embed(tgt))\n        out = self.decoder(\n            emb, memory,\n            tgt_mask                = causal,\n            tgt_key_padding_mask    = tgt_pad,\n            memory_key_padding_mask = memory_mask,\n        )\n        return self.proj(out)\n\n    @torch.no_grad()\n    def generate(\n        self,\n        memory      : torch.Tensor,\n        memory_mask : Optional[torch.Tensor] = None,\n        max_len     : int = 120,\n        bos_id      : int = CharTokenizer.BOS_ID,\n        eos_id      : int = CharTokenizer.EOS_ID,\n    ) -> torch.Tensor:\n        \"\"\"Greedy autoregressive decoding. Returns (B, ≤ max_len+1) token IDs.\"\"\"\n        B      = memory.size(0)\n        device = memory.device\n        tokens = torch.full((B, 1), bos_id, dtype=torch.long, device=device)\n        done   = torch.zeros(B, dtype=torch.bool, device=device)\n\n        for _ in range(max_len):\n            logits     = self.forward(tokens, memory, memory_mask)\n            next_tok   = logits[:, -1].argmax(dim=-1, keepdim=True)  # (B, 1)\n            tokens     = torch.cat([tokens, next_tok], dim=1)\n            done       = done | (next_tok.squeeze(-1) == eos_id)\n            if done.all():\n                break\n\n        return tokens\n\n\nprint('AttentionDecoder defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:33.869017Z","iopub.execute_input":"2026-03-07T11:45:33.869257Z","iopub.status.idle":"2026-03-07T11:45:33.887442Z","shell.execute_reply.started":"2026-03-07T11:45:33.869236Z","shell.execute_reply":"2026-03-07T11:45:33.886938Z"}},"outputs":[],"execution_count":null},{"id":"e716c2ab-0e7f-4329-acd5-47d1a5ece866","cell_type":"markdown","source":"## Cell 8 — Full DSD-NLA v2 Model","metadata":{}},{"id":"767a7744-92fe-47e1-87c4-56ec5daa7376","cell_type":"code","source":"\"\"\"\nCELL 8 — DSD-NLA v2\n\"\"\"\nclass DSDNLA_v2(nn.Module):\n    \"\"\"\n    DSD-NLA v2: Direct Spectral Decoding Neural Language Architecture v2\n\n    Components:\n      1. ConformerEncoder   → z  (B, T//2, D)\n      2. CTC Head           → ctc_log_probs  (B, T//2, V)\n      3. AttentionDecoder   → dec_logits  (B, T_text, V)\n\n    Joint loss = ctc_weight * CTC + (1 - ctc_weight) * CE\n    \"\"\"\n    def __init__(self, cfg: Config, vocab_size: int):\n        super().__init__()\n        self.cfg = cfg\n\n        self.encoder = ConformerEncoder(\n            n_channels   = cfg.n_channels,\n            d_model      = cfg.d_model,\n            n_layers     = cfg.n_encoder_layers,\n            n_heads      = cfg.n_heads,\n            ff_expansion = cfg.ff_expansion,\n            conv_kernel  = cfg.conv_kernel,\n            dropout      = cfg.dropout,\n            drop_path    = cfg.stochastic_depth_prob,\n        )\n        self.ctc_head = nn.Sequential(\n            nn.LayerNorm(cfg.d_model),\n            nn.Linear(cfg.d_model, vocab_size),\n        )\n        self.decoder = AttentionDecoder(\n            vocab_size   = vocab_size,\n            d_model      = cfg.d_model,\n            n_layers     = cfg.n_decoder_layers,\n            n_heads      = cfg.n_heads,\n            ff_expansion = cfg.ff_expansion,\n            dropout      = cfg.dropout,\n        )\n\n    def forward(\n        self,\n        neural        : torch.Tensor,\n        neural_mask   : Optional[torch.Tensor] = None,\n        target_tokens : Optional[torch.Tensor] = None,\n    ) -> Dict[str, torch.Tensor]:\n        z, enc_mask = self.encoder(neural, neural_mask)\n        ctc_log     = F.log_softmax(self.ctc_head(z), dim=-1)\n\n        dec_logits = None\n        if target_tokens is not None:\n            dec_logits = self.decoder(target_tokens, z, enc_mask)\n\n        return {\n            'ctc_log_probs' : ctc_log,\n            'enc_mask'      : enc_mask,\n            'z'             : z,\n            'dec_logits'    : dec_logits,\n        }\n\n    @torch.no_grad()\n    def inference(\n        self,\n        neural      : torch.Tensor,\n        neural_mask : Optional[torch.Tensor] = None,\n        use_ctc     : bool = True,\n        max_len     : int  = 120,\n    ) -> Dict[str, Any]:\n        z, enc_mask = self.encoder(neural, neural_mask)\n        ctc_log     = F.log_softmax(self.ctc_head(z), dim=-1)\n\n        tokens = None\n        if not use_ctc:\n            tokens = self.decoder.generate(z, enc_mask, max_len=max_len)\n\n        return {'ctc_log_probs': ctc_log, 'tokens': tokens, 'enc_mask': enc_mask}\n\n\n# ── Smoke test ─────────────────────────────────────────────────────\nmodel = DSDNLA_v2(CFG, tokenizer.vocab_size).to(CFG.device)\nn_params = sum(p.numel() for p in model.parameters())\nprint(f'Model parameters : {n_params:,}')\n\n_x   = torch.randn(2, 200, 512, device=CFG.device)\n_tgt = torch.randint(4, 30, (2, 20), device=CFG.device)\n_out = model(_x, target_tokens=_tgt)\nprint(f\"CTC log-probs : {_out['ctc_log_probs'].shape}\")\nprint(f\"Dec logits    : {_out['dec_logits'].shape}\")\nprint('Smoke test passed.')\ndel _x, _tgt, _out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:33.889175Z","iopub.execute_input":"2026-03-07T11:45:33.889386Z","iopub.status.idle":"2026-03-07T11:45:34.516447Z","shell.execute_reply.started":"2026-03-07T11:45:33.889367Z","shell.execute_reply":"2026-03-07T11:45:34.515858Z"}},"outputs":[],"execution_count":null},{"id":"898b2cb8-41d8-45b6-ab57-9f0dbaf7be26","cell_type":"markdown","source":"## Cell 9 — Hybrid CTC + Cross-Entropy Loss\n\n### Why CTC?\nCTC is **alignment-free**: it does not require knowing which encoder frame\ncorresponds to which output character. This is important for brain signals\nbecause neural activity is not precisely aligned to speech events.\n\n### Key CTC Details\n- `F.ctc_loss` expects log-probs in shape `(T, B, V)` — time-major\n- Targets must be a **flat 1-D concatenation** of all (non-padded) targets\n- `input_lengths` = number of valid encoder frames (from `enc_mask`)\n- `zero_infinity=True` suppresses NaN from impossible CTC alignments\n\n### Fix: `input_lengths` computation\nPreviously: `neural_lens // 2` — wrong when Conv padding creates off-by-one.\nNow: count non-PAD frames directly from `enc_mask`.","metadata":{}},{"id":"dc91bf93-feae-4b0d-9a74-5e6df8e15015","cell_type":"code","source":"\"\"\"\nCELL 9 — Hybrid CTC + Cross-Entropy Loss\n\"\"\"\nclass HybridCTCCELoss(nn.Module):\n    def __init__(\n        self,\n        ctc_weight     : float = 0.3,\n        label_smoothing: float = 0.1,\n        ignore_index   : int   = CharTokenizer.PAD_ID,\n    ):\n        super().__init__()\n        self.ctc_weight = ctc_weight\n        self.ctc = nn.CTCLoss(\n            blank=CharTokenizer.BLANK_ID,\n            reduction='mean',\n            zero_infinity=True,   # avoids NaN from impossible alignments\n        )\n        self.ce = nn.CrossEntropyLoss(\n            ignore_index=ignore_index,\n            label_smoothing=label_smoothing,\n        )\n\n    def forward(\n        self,\n        ctc_log_probs : torch.Tensor,              # (B, L, V)\n        enc_mask      : Optional[torch.Tensor],    # (B, L) True=PAD\n        ctc_targets   : torch.Tensor,              # (B, T_ctc) padded\n        ctc_lens      : torch.Tensor,              # (B,) actual CTC lengths\n        dec_logits    : Optional[torch.Tensor],    # (B, T_text, V)\n        token_targets : Optional[torch.Tensor],    # (B, T_text)\n    ) -> Dict[str, torch.Tensor]:\n\n        device = ctc_log_probs.device\n        B, L, V = ctc_log_probs.shape\n\n        # ── CTC input_lengths from enc_mask ──────────────────────────\n        # enc_mask: True=PAD → valid frames = ~enc_mask\n        if enc_mask is not None:\n            input_lengths = (~enc_mask).sum(dim=1).long()  # (B,)\n            input_lengths = input_lengths.clamp(min=1, max=L)\n        else:\n            input_lengths = torch.full((B,), L, dtype=torch.long, device=device)\n\n        # ctc_loss needs (T, B, V) — time-major\n        log_p_t = ctc_log_probs.transpose(0, 1).contiguous()  # (L, B, V)\n\n        # Flatten targets: concat each example's non-padded CTC labels\n        tgt_flat = torch.cat(\n            [ctc_targets[i, : ctc_lens[i].item()] for i in range(B)]\n        )\n\n        loss_ctc = self.ctc(log_p_t, tgt_flat, input_lengths, ctc_lens)\n\n        # ── Cross-entropy (attention decoder) ────────────────────────\n        loss_ce = torch.tensor(0.0, device=device)\n        if dec_logits is not None and token_targets is not None:\n            # Teacher forcing: predict token[t+1] given token[t]\n            B2, T, V2 = dec_logits.shape\n            loss_ce = self.ce(\n                dec_logits[:, :-1].reshape(-1, V2),\n                token_targets[:, 1:].reshape(-1),\n            )\n\n        total = self.ctc_weight * loss_ctc + (1.0 - self.ctc_weight) * loss_ce\n        return {'total': total, 'ctc': loss_ctc, 'ce': loss_ce}\n\n\nprint('HybridCTCCELoss defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:34.517286Z","iopub.execute_input":"2026-03-07T11:45:34.517619Z","iopub.status.idle":"2026-03-07T11:45:34.526624Z","shell.execute_reply.started":"2026-03-07T11:45:34.517591Z","shell.execute_reply":"2026-03-07T11:45:34.525996Z"}},"outputs":[],"execution_count":null},{"id":"11a8f3c8-2edd-4f85-a291-7e2391f67804","cell_type":"markdown","source":"## Cell 10 — WER Metric","metadata":{}},{"id":"e31ad73c-8079-48dd-b24a-ad580a0cf8b6","cell_type":"code","source":"\"\"\"\nCELL 10 — Word Error Rate\n\nWER = (S + D + I) / N\n  S = word substitutions\n  D = word deletions\n  I = word insertions\n  N = number of words in reference\n\nThis is the standard competition metric.\nLower is better; 0.0 = perfect.\n\"\"\"\ndef wer_single(reference: str, hypothesis: str) -> float:\n    ref = reference.strip().lower().split()\n    hyp = hypothesis.strip().lower().split()\n    if len(ref) == 0:\n        return 0.0 if len(hyp) == 0 else 1.0\n\n    R, H = len(ref), len(hyp)\n    # DP table: dp[i][j] = edit distance between ref[:i] and hyp[:j]\n    dp = [[0] * (H + 1) for _ in range(R + 1)]\n    for i in range(R + 1):\n        dp[i][0] = i\n    for j in range(H + 1):\n        dp[0][j] = j\n    for i in range(1, R + 1):\n        for j in range(1, H + 1):\n            cost = 0 if ref[i - 1] == hyp[j - 1] else 1\n            dp[i][j] = min(\n                dp[i - 1][j] + 1,        # deletion\n                dp[i][j - 1] + 1,        # insertion\n                dp[i - 1][j - 1] + cost  # substitution\n            )\n    return dp[R][H] / R\n\n\ndef corpus_wer(references: List[str], hypotheses: List[str]) -> float:\n    \"\"\"Macro-average WER over a corpus.\"\"\"\n    assert len(references) == len(hypotheses)\n    if len(references) == 0:\n        return 0.0\n    return float(np.mean([wer_single(r, h) for r, h in zip(references, hypotheses)]))\n\n\n# Unit tests\nassert wer_single('hello world', 'hello world') == 0.0\nassert wer_single('hello world', 'hello') == 0.5\nassert wer_single('', '') == 0.0\nprint(f'WER tests passed.')\nprint(f\"  wer('the cat sat on the mat', 'the cat sat on mat') = {wer_single('the cat sat on the mat', 'the cat sat on mat'):.3f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:34.527409Z","iopub.execute_input":"2026-03-07T11:45:34.527695Z","iopub.status.idle":"2026-03-07T11:45:34.552195Z","shell.execute_reply.started":"2026-03-07T11:45:34.527666Z","shell.execute_reply":"2026-03-07T11:45:34.551615Z"}},"outputs":[],"execution_count":null},{"id":"512ac76a-c726-43b8-b896-c27fd5b0e341","cell_type":"markdown","source":"## Cell 11 — Trainer","metadata":{}},{"id":"f72919b5-3ee3-43fa-b6a8-3fc807bf79e9","cell_type":"code","source":"\"\"\"\nCELL 11 — Trainer\n\"\"\"\nclass Trainer:\n    def __init__(self, model: DSDNLA_v2, tok: CharTokenizer, cfg: Config):\n        self.model = model.to(cfg.device)\n        self.tok   = tok\n        self.cfg   = cfg\n\n        self.criterion = HybridCTCCELoss(cfg.ctc_weight, cfg.label_smoothing)\n\n        self.optimizer = optim.AdamW(\n            model.parameters(),\n            lr           = cfg.learning_rate,\n            weight_decay = cfg.weight_decay,\n            betas        = (0.9, 0.98),\n            eps          = 1e-8,\n        )\n        self.scheduler: Optional[optim.lr_scheduler.OneCycleLR] = None\n        self.scaler = (\n            torch.amp.GradScaler('cuda')\n            if cfg.use_amp and cfg.device == 'cuda'\n            else None\n        )\n\n        self.global_step  = 0\n        self.best_val_wer = float('inf')\n        self.patience_cnt = 0\n        self.history: Dict[str, List[float]] = {\n            'train_loss': [], 'val_loss': [], 'val_wer': []\n        }\n\n    def _build_scheduler(self, steps_per_epoch: int):\n        total = max(1, steps_per_epoch * self.cfg.num_epochs)\n        self.scheduler = optim.lr_scheduler.OneCycleLR(\n            self.optimizer,\n            max_lr          = self.cfg.learning_rate,\n            total_steps     = total,\n            pct_start       = self.cfg.warmup_frac,\n            anneal_strategy = 'cos',\n        )\n\n    def _compute_loss(\n        self,\n        batch: Dict[str, Any],\n    ) -> Tuple[Dict[str, torch.Tensor], Dict[str, torch.Tensor]]:\n        \"\"\"Shared forward + loss for train and val. Returns (losses, model_out).\"\"\"\n        neural      = batch['neural'].to(self.cfg.device)\n        neural_mask = batch['neural_mask'].to(self.cfg.device)\n        tokens      = batch['tokens'].to(self.cfg.device)\n        ctc_tgt     = batch['ctc_tgt'].to(self.cfg.device)\n        ctc_lens    = batch['ctc_lens'].to(self.cfg.device)\n\n        use_amp = (self.scaler is not None)\n        with torch.amp.autocast('cuda', enabled=use_amp):\n            out    = self.model(neural, neural_mask, target_tokens=tokens)\n            losses = self.criterion(\n                out['ctc_log_probs'],\n                out['enc_mask'],\n                ctc_tgt, ctc_lens,\n                out['dec_logits'], tokens,\n            )\n        return losses, out\n\n    def train_epoch(self, loader: DataLoader, epoch: int) -> float:\n        self.model.train()\n        epoch_losses: List[float] = []\n        pbar = tqdm(loader, desc=f'Epoch {epoch:02d} [train]')\n\n        for batch in pbar:\n            losses, _ = self._compute_loss(batch)\n            loss      = losses['total']\n\n            self.optimizer.zero_grad(set_to_none=True)\n            if self.scaler is not None:\n                self.scaler.scale(loss).backward()\n                self.scaler.unscale_(self.optimizer)\n                nn.utils.clip_grad_norm_(self.model.parameters(), self.cfg.gradient_clip)\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n            else:\n                loss.backward()\n                nn.utils.clip_grad_norm_(self.model.parameters(), self.cfg.gradient_clip)\n                self.optimizer.step()\n\n            if self.scheduler is not None:\n                self.scheduler.step()\n\n            epoch_losses.append(loss.item())\n            pbar.set_postfix({\n                'loss': f\"{loss.item():.4f}\",\n                'ctc' : f\"{losses['ctc'].item():.4f}\",\n                'ce'  : f\"{losses['ce'].item():.4f}\",\n            })\n            self.global_step += 1\n\n        return float(np.mean(epoch_losses))\n\n    @torch.no_grad()\n    def validate(self, loader: DataLoader) -> Tuple[float, float]:\n        self.model.eval()\n        val_losses: List[float] = []\n        refs: List[str] = []\n        hyps: List[str] = []\n\n        for batch in tqdm(loader, desc='Validation', leave=False):\n            losses, out = self._compute_loss(batch)\n            val_losses.append(losses['total'].item())\n            preds = self.tok.ctc_decode_greedy(out['ctc_log_probs'])\n            refs.extend(batch['texts'])\n            hyps.extend(preds)\n\n        return float(np.mean(val_losses)), corpus_wer(refs, hyps)\n\n    def save(self, epoch: int):\n        Path(self.cfg.checkpoint_path).parent.mkdir(parents=True, exist_ok=True)\n        torch.save({\n            'epoch'     : epoch,\n            'step'      : self.global_step,\n            'model'     : self.model.state_dict(),\n            'optimizer' : self.optimizer.state_dict(),\n            'scheduler' : self.scheduler.state_dict() if self.scheduler else None,\n            'best_wer'  : self.best_val_wer,\n            'history'   : self.history,\n        }, self.cfg.checkpoint_path)\n        print(f'  ✓ Checkpoint saved → {self.cfg.checkpoint_path}  (WER={self.best_val_wer:.4f})')\n\n    def train(self, train_loader: DataLoader, val_loader: DataLoader):\n        print('=' * 65)\n        print('  DSD-NLA v2 — Training')\n        print(f'  Device : {self.cfg.device}')\n        print(f'  Epochs : {self.cfg.num_epochs}   Patience : {self.cfg.patience}')\n        print(f'  CTC α  : {self.cfg.ctc_weight}   Batch : {self.cfg.batch_size}')\n        print('=' * 65)\n\n        self._build_scheduler(len(train_loader))\n\n        for epoch in range(self.cfg.num_epochs):\n            tr_loss           = self.train_epoch(train_loader, epoch)\n            val_loss, val_wer = self.validate(val_loader)\n\n            self.history['train_loss'].append(tr_loss)\n            self.history['val_loss'].append(val_loss)\n            self.history['val_wer'].append(val_wer)\n\n            lr_now = self.optimizer.param_groups[0]['lr']\n            print(f'\\nEpoch {epoch:02d} | '\n                  f'TrainLoss={tr_loss:.4f} | '\n                  f'ValLoss={val_loss:.4f} | '\n                  f'ValWER={val_wer:.4f} | '\n                  f'LR={lr_now:.2e}')\n\n            if val_wer < self.best_val_wer:\n                self.best_val_wer = val_wer\n                self.patience_cnt = 0\n                self.save(epoch)\n            else:\n                self.patience_cnt += 1\n                print(f'  No improvement ({self.patience_cnt}/{self.cfg.patience})')\n                if self.patience_cnt >= self.cfg.patience:\n                    print(f'\\n  Early stopping at epoch {epoch}.')\n                    break\n\n        print('\\n' + '=' * 65)\n        print(f'  Training complete.  Best Val WER = {self.best_val_wer:.4f}')\n        print('=' * 65)\n\n\nprint('Trainer defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:34.553311Z","iopub.execute_input":"2026-03-07T11:45:34.553664Z","iopub.status.idle":"2026-03-07T11:45:34.576931Z","shell.execute_reply.started":"2026-03-07T11:45:34.553641Z","shell.execute_reply":"2026-03-07T11:45:34.576329Z"}},"outputs":[],"execution_count":null},{"id":"e14d4718-94dd-45ab-b497-59433999bf1e","cell_type":"markdown","source":"## Cell 12 — Data Loading & Training","metadata":{}},{"id":"e7392d81-e31b-4eb7-b369-395c844e70f5","cell_type":"code","source":"\"\"\"\nCELL 12 — Build DataLoaders & Run Training\n\nSet RUN_TRAINING = False to skip training and jump straight to inference\n(useful when reloading a pre-trained checkpoint).\n\"\"\"\nRUN_TRAINING = True\n\n# Global references so downstream cells can access them\ntrainer          = None\ntrain_loader_ref = None\nval_loader_ref   = None\n\n\ndef build_dataloaders(\n    cfg : Config,\n    tok : CharTokenizer,\n) -> Tuple[DataLoader, DataLoader]:\n    train_files = sorted(glob.glob(f'{cfg.data_dir}/t15.*/data_train.hdf5'))\n    val_files   = sorted(glob.glob(f'{cfg.data_dir}/t15.*/data_val.hdf5'))\n\n    print(f'Train files : {len(train_files)}')\n    print(f'Val   files : {len(val_files)}')\n\n    # ── Helpful error if data is not mounted ─────────────────────────\n    if len(train_files) == 0:\n        raise FileNotFoundError(\n            f'No training HDF5 files found at:\\n'\n            f'  {cfg.data_dir}/t15.*/data_train.hdf5\\n\\n'\n            f'Checklist:\\n'\n            f'  1. Is the Brain-to-Text 25 dataset added to this notebook?\\n'\n            f'     (Kaggle → Data → Add Data → search \"brain-to-text-25\")\\n'\n            f'  2. Does the data_dir path in Config match the mounted location?\\n'\n            f'     Current: {cfg.data_dir}\\n'\n            f'  3. Run Cell 3 (Data Explorer) to verify the folder structure.'\n        )\n\n    train_ds = BrainToTextDataset(\n        train_files, tok, mode='train', augment=True, max_neural=cfg.max_neural_len\n    )\n    val_ds = BrainToTextDataset(\n        val_files, tok, mode='val', augment=False, max_neural=cfg.max_neural_len\n    )\n\n    train_loader = DataLoader(\n        train_ds,\n        batch_size  = cfg.batch_size,\n        shuffle     = True,\n        collate_fn  = collate_fn,\n        num_workers = 2,\n        pin_memory  = True,\n        drop_last   = True,   # avoids CTC issues with very small last batch\n    )\n    val_loader = DataLoader(\n        val_ds,\n        batch_size  = cfg.batch_size,\n        shuffle     = False,\n        collate_fn  = collate_fn,\n        num_workers = 2,\n        pin_memory  = True,\n    )\n    return train_loader, val_loader\n\n\nif RUN_TRAINING:\n    train_loader_ref, val_loader_ref = build_dataloaders(CFG, tokenizer)\n\n    model   = DSDNLA_v2(CFG, tokenizer.vocab_size)\n    trainer = Trainer(model, tokenizer, CFG)\n    trainer.train(train_loader_ref, val_loader_ref)\n\n    print('\\nFiles in /kaggle/working:')\n    for f in sorted(Path('/kaggle/working').iterdir()):\n        print(f'  {f.name}')\nelse:\n    print(f'Skipping training.  Checkpoint expected at: {CFG.checkpoint_path}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T11:45:34.577891Z","iopub.execute_input":"2026-03-07T11:45:34.578221Z"}},"outputs":[],"execution_count":null},{"id":"e50199e8-8553-4116-9833-96ddd1efbbed","cell_type":"markdown","source":"## Cell 13 — Training Curves","metadata":{}},{"id":"8f6ce5fe-f64d-4865-82bf-4e9f96349ea3","cell_type":"code","source":"\"\"\"\nCELL 13 — Plot training curves\n\"\"\"\nimport matplotlib.pyplot as plt\n\n\ndef plot_training_curves(history: Dict[str, List[float]], save_path: str = None):\n    if not history or len(history.get('train_loss', [])) == 0:\n        print('No training history to plot.')\n        return\n\n    epochs = range(len(history['train_loss']))\n    fig, axes = plt.subplots(1, 2, figsize=(14, 4))\n\n    axes[0].plot(epochs, history['train_loss'], 'o-', markersize=3,\n                 color='steelblue', label='Train Loss')\n    axes[0].plot(epochs, history['val_loss'],   's-', markersize=3,\n                 color='coral',    label='Val Loss')\n    axes[0].set_xlabel('Epoch'); axes[0].set_ylabel('Hybrid Loss (CTC + CE)')\n    axes[0].set_title('Training & Validation Loss', fontweight='bold')\n    axes[0].legend(); axes[0].grid(alpha=0.3)\n\n    axes[1].plot(epochs, history['val_wer'], 'D-', markersize=3, color='darkorange')\n    best_ep  = int(np.argmin(history['val_wer']))\n    best_wer = history['val_wer'][best_ep]\n    axes[1].axvline(best_ep, color='gray', linestyle='--', alpha=0.6,\n                    label=f'Best (ep {best_ep}, WER={best_wer:.4f})')\n    axes[1].set_xlabel('Epoch'); axes[1].set_ylabel('Val WER (CTC greedy)')\n    axes[1].set_title(f'Validation WER  [best={best_wer:.4f}]', fontweight='bold')\n    axes[1].legend(); axes[1].grid(alpha=0.3)\n\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path, dpi=150, bbox_inches='tight')\n        print(f'Curves saved → {save_path}')\n    plt.show()\n\n\nif trainer is not None:\n    plot_training_curves(trainer.history, '/kaggle/working/training_curves.png')\nelse:\n    print('trainer not available (RUN_TRAINING=False). Skipping plot.')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}