{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.0"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":106809,"databundleVersionId":13056355}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# DSD-NLA v2: Neural-to-Text Decoding — Research Edition\n\n> **Competition:** Brain-to-Text '25 (Kaggle)  \n> **Task:** Decode intracortical neural activity → 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 presents **DSD-NLA v2**, an improved end-to-end brain-to-text decoding system.\n\nKey improvements over v1:\n\n| Component | v1 (baseline) | v2 (this notebook) |\n|---|---|---|\n| Encoder | Plain Transformer | **Conformer** (conv-augmented attention) |\n| Loss | CE only (teacher forcing) | **Joint CTC + Label-Smoothed CE** |\n| Regularisation | Dropout | Dropout + **Stochastic Depth** |\n| Decoding | Greedy AR | **CTC greedy** + Attention greedy |\n| Augmentation | Temporal mask, noise | + **Channel dropout**, time-shift |\n| Scheduler | OneCycleLR | OneCycleLR + **cosine warmup** |\n\n---\n\n## Architecture\n\n```\nNeural Signal  (B, T, 512)\n      │\n      ▼\n Conv1d Subsampling (stride=2)  →  (B, T//2, D)\n      │\n      ▼\n Conformer Encoder  (8 × ConformerBlock)\n      │\n      ├──► CTC Head  →  L_CTC   (alignment-free)\n      │\n      ▼\n Attention Decoder  (6 × TransformerDecoderLayer)\n      │\n      ▼\n L_total = 0.3 × L_CTC  +  0.7 × L_CE\n```\n","metadata":{}},{"cell_type":"markdown","source":"## Cell 1: Imports & Environment","metadata":{}},{"cell_type":"code","source":"import os, re, math, string, glob\nfrom pathlib import Path\nfrom typing import List, Dict, Tuple, Optional, Any\nfrom dataclasses import dataclass, field\n\nimport numpy as np\nimport h5py\nimport pandas as pd\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\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-03T18:42:14.632982Z","iopub.execute_input":"2026-03-03T18:42:14.63334Z","iopub.status.idle":"2026-03-03T18:42:19.362625Z","shell.execute_reply.started":"2026-03-03T18:42:14.633283Z","shell.execute_reply":"2026-03-03T18:42:19.361956Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 2: Configuration","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass Config:\n    # Paths\n    data_dir: str        = \"/kaggle/input/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    # Data\n    n_channels: int = 512\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\n    ff_expansion: int        = 4\n    dropout: float           = 0.15\n    stochastic_depth_prob: float = 0.10\n    ctc_weight: float        = 0.3\n    label_smoothing: float   = 0.1\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\n    patience: int        = 8\n    use_amp: bool        = True\n    # Inference\n    ctc_decode: bool = True\n    # Device\n    device: str = field(default_factory=lambda: \"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nCFG = Config()\ntorch.manual_seed(CFG.seed)\nnp.random.seed(CFG.seed)\nprint(\"\\n=== Config ===\")\nfor k, v in CFG.__dict__.items():\n    print(f\"  {k:<32} {v}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:36.029996Z","iopub.execute_input":"2026-03-03T18:42:36.030482Z","iopub.status.idle":"2026-03-03T18:42:36.042653Z","shell.execute_reply.started":"2026-03-03T18:42:36.030453Z","shell.execute_reply":"2026-03-03T18:42:36.041955Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 3: Character Tokenizer","metadata":{}},{"cell_type":"code","source":"class CharTokenizer:\n    BLANK_ID = 0   # CTC blank\n    PAD_ID   = 1\n    BOS_ID   = 2\n    EOS_ID   = 3\n\n    def __init__(self):\n        self.chars = [\"<blank>\", \"<PAD>\", \"<BOS>\", \"<EOS>\"]\n        self.chars += list(string.ascii_lowercase)\n        self.chars += [\" \", \"'\", \".\", \",\", \"!\", \"?\"]\n        self.char2id = {c: i for i, c in enumerate(self.chars)}\n        self.id2char  = {i: c for i, c in enumerate(self.chars)}\n        self.vocab_size = len(self.chars)\n\n    def encode(self, text: str, add_bos_eos: bool = True) -> torch.Tensor:\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        return self.encode(text, add_bos_eos=False)\n\n    def decode(self, ids, skip_special: bool = True) -> str:\n        if isinstance(ids, torch.Tensor):\n            ids = ids.cpu().tolist()\n        out = []\n        for i in ids:\n            if skip_special:\n                if i == self.EOS_ID: break\n                if i in (self.BLANK_ID, self.PAD_ID, self.BOS_ID): continue\n            out.append(self.id2char.get(i, \" \"))\n        return \"\".join(out)\n\n    def ctc_greedy_decode(self, log_probs: torch.Tensor) -> str:\n        ids = log_probs.argmax(dim=-1).cpu().tolist()\n        collapsed = [ids[0]] + [ids[i] for i in range(1, len(ids)) if ids[i] != ids[i-1]]\n        return self.decode([i for i in collapsed if i != self.BLANK_ID])\n\nTOKENIZER = CharTokenizer()\nprint(f\"Vocab size: {TOKENIZER.vocab_size}\")\n_enc = TOKENIZER.encode(\"hello world\")\nprint(f\"encode('hello world') → {_enc.tolist()}\")\nprint(f\"decode back → '{TOKENIZER.decode(_enc)}'\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:39.778229Z","iopub.execute_input":"2026-03-03T18:42:39.778568Z","iopub.status.idle":"2026-03-03T18:42:39.804136Z","shell.execute_reply.started":"2026-03-03T18:42:39.778542Z","shell.execute_reply":"2026-03-03T18:42:39.803586Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 4: Conformer Building Blocks\n\n> Reference: Gulati et al. 2020 — *Conformer: Convolution-augmented Transformer for Speech Recognition* (https://arxiv.org/abs/2005.08100)","metadata":{}},{"cell_type":"code","source":"# ── Positional Encoding ──────────────────────────────────────────────────\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, dropout=0.1, max_len=5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n        pe  = torch.zeros(max_len, d_model)\n        pos = torch.arange(0, max_len).float().unsqueeze(1)\n        div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(pos * div)\n        pe[:, 1::2] = torch.cos(pos * div)\n        self.register_buffer(\"pe\", pe.unsqueeze(0))\n    def forward(self, x):\n        return self.dropout(x + self.pe[:, :x.size(1)])\n\n\n# ── Swish ────────────────────────────────────────────────────────────────\nclass Swish(nn.Module):\n    def forward(self, x): return x * torch.sigmoid(x)\n\n\n# ── Feed-Forward Module (half-step residual) ─────────────────────────────\nclass ConformerFFN(nn.Module):\n    def __init__(self, d_model, expansion=4, dropout=0.1):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.LayerNorm(d_model),\n            nn.Linear(d_model, d_model * expansion),\n            Swish(),\n            nn.Dropout(dropout),\n            nn.Linear(d_model * expansion, d_model),\n            nn.Dropout(dropout),\n        )\n    def forward(self, x): return self.net(x)\n\n\n# ── Depthwise Convolution Module ─────────────────────────────────────────\nclass ConformerConvModule(nn.Module):\n    def __init__(self, d_model, kernel_size=31, dropout=0.1):\n        super().__init__()\n        assert (kernel_size - 1) % 2 == 0\n        self.ln   = nn.LayerNorm(d_model)\n        self.pw1  = nn.Conv1d(d_model, 2*d_model, 1)\n        self.glu  = nn.GLU(dim=1)\n        self.dw   = nn.Conv1d(d_model, d_model, kernel_size,\n                              padding=(kernel_size-1)//2, groups=d_model)\n        self.bn   = nn.BatchNorm1d(d_model)\n        self.act  = Swish()\n        self.pw2  = nn.Conv1d(d_model, d_model, 1)\n        self.drop = nn.Dropout(dropout)\n    def forward(self, x):          # x: (B, T, D)\n        x = self.ln(x).transpose(1, 2)\n        x = self.glu(self.pw1(x))\n        x = self.act(self.bn(self.dw(x)))\n        return self.drop(self.pw2(x)).transpose(1, 2)\n\n\n# ── Conformer Block ───────────────────────────────────────────────────────\nclass ConformerBlock(nn.Module):\n    def __init__(self, d_model, n_heads, conv_kernel, ff_expansion, dropout, layer_drop=0.0):\n        super().__init__()\n        self.ff1      = ConformerFFN(d_model, ff_expansion, dropout)\n        self.ln_att   = nn.LayerNorm(d_model)\n        self.att      = nn.MultiheadAttention(d_model, n_heads, dropout=dropout, batch_first=True)\n        self.att_drop = nn.Dropout(dropout)\n        self.conv     = ConformerConvModule(d_model, conv_kernel, dropout)\n        self.ff2      = ConformerFFN(d_model, ff_expansion, dropout)\n        self.ln       = nn.LayerNorm(d_model)\n        self.layer_drop = layer_drop\n\n    def forward(self, x, key_padding_mask=None):\n        if self.training and self.layer_drop > 0 and torch.rand(1).item() < self.layer_drop:\n            return x\n        x  = x + 0.5 * self.ff1(x)\n        x_ = self.ln_att(x)\n        x_, _ = self.att(x_, x_, x_, key_padding_mask=key_padding_mask)\n        x  = x + self.att_drop(x_)\n        x  = x + self.conv(x)\n        x  = x + 0.5 * self.ff2(x)\n        return self.ln(x)\n\nprint(\"Conformer building blocks ✓\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:42.855928Z","iopub.execute_input":"2026-03-03T18:42:42.856478Z","iopub.status.idle":"2026-03-03T18:42:42.869886Z","shell.execute_reply.started":"2026-03-03T18:42:42.856448Z","shell.execute_reply":"2026-03-03T18:42:42.869173Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 5: Full DSD-NLA v2 Model","metadata":{}},{"cell_type":"code","source":"# ── Neural Encoder (Conformer) ────────────────────────────────────────────\nclass NeuralEncoder(nn.Module):\n    def __init__(self, cfg: Config, vocab_size: int):\n        super().__init__()\n        D = cfg.d_model\n        self.input_proj = nn.Sequential(\n            nn.Conv1d(cfg.n_channels, D, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(D), nn.GELU(), nn.Dropout(cfg.dropout),\n        )\n        self.pos_enc = PositionalEncoding(D, cfg.dropout)\n        n = cfg.n_encoder_layers\n        drops = [i / n * cfg.stochastic_depth_prob for i in range(n)]\n        self.conformer = nn.ModuleList([\n            ConformerBlock(D, cfg.n_heads, cfg.conv_kernel,\n                           cfg.ff_expansion, cfg.dropout, drops[i])\n            for i in range(n)\n        ])\n        self.ln = nn.LayerNorm(D)\n        self.ctc_head = nn.Linear(D, vocab_size)\n\n    def forward(self, x, pad_mask=None):\n        x = self.input_proj(x.transpose(1, 2)).transpose(1, 2)\n        enc_mask = None\n        if pad_mask is not None:\n            enc_mask = pad_mask[:, ::2]\n            L = x.size(1)\n            if enc_mask.size(1) > L:   enc_mask = enc_mask[:, :L]\n            elif enc_mask.size(1) < L:\n                pad = torch.ones(enc_mask.size(0), L - enc_mask.size(1),\n                                 dtype=torch.bool, device=enc_mask.device)\n                enc_mask = torch.cat([enc_mask, pad], dim=1)\n        x = self.pos_enc(x)\n        for blk in self.conformer:\n            x = blk(x, key_padding_mask=enc_mask)\n        z = self.ln(x)\n        return z, self.ctc_head(z), enc_mask\n\n\n# ── Attention Text Decoder ────────────────────────────────────────────────\nclass TextDecoder(nn.Module):\n    def __init__(self, cfg: Config, vocab_size: int):\n        super().__init__()\n        D = cfg.d_model\n        self.embed   = nn.Embedding(vocab_size, D)\n        self.pos_enc = PositionalEncoding(D, cfg.dropout)\n        dec_layer = nn.TransformerDecoderLayer(\n            d_model=D, nhead=cfg.n_heads,\n            dim_feedforward=D*cfg.ff_expansion,\n            dropout=cfg.dropout, activation=\"gelu\",\n            batch_first=True, norm_first=True,\n        )\n        self.decoder = nn.TransformerDecoder(dec_layer, num_layers=cfg.n_decoder_layers)\n        self.proj = nn.Linear(D, vocab_size)\n\n    def forward(self, z_enc, tgt_tokens, memory_key_padding_mask=None):\n        T = tgt_tokens.size(1)\n        causal = nn.Transformer.generate_square_subsequent_mask(T, device=z_enc.device)\n        tgt = self.pos_enc(self.embed(tgt_tokens))\n        out = self.decoder(tgt, z_enc, tgt_mask=causal,\n                           memory_key_padding_mask=memory_key_padding_mask)\n        return self.proj(out)\n\n    @torch.no_grad()\n    def generate(self, z_enc, max_len=100, bos_id=2, eos_id=3, mem_mask=None):\n        B, device = z_enc.size(0), z_enc.device\n        gen = torch.full((B, 1), bos_id, dtype=torch.long, device=device)\n        for _ in range(max_len):\n            logits = self.forward(z_enc, gen, mem_mask)\n            nxt = logits[:, -1].argmax(dim=-1, keepdim=True)\n            gen = torch.cat([gen, nxt], dim=1)\n            if (nxt == eos_id).all(): break\n        return gen\n\n\n# ── DSD-NLA v2 (full model) ───────────────────────────────────────────────\nclass DSDNLAv2(nn.Module):\n    def __init__(self, cfg: Config, vocab_size: int):\n        super().__init__()\n        self.encoder = NeuralEncoder(cfg, vocab_size)\n        self.decoder = TextDecoder(cfg, vocab_size)\n\n    def forward(self, neural, tgt_tokens=None, neural_mask=None):\n        z, ctc_logits, enc_mask = self.encoder(neural, pad_mask=neural_mask)\n        out = {\"ctc_logits\": ctc_logits, \"enc_mask\": enc_mask, \"z_enc\": z}\n        if tgt_tokens is not None:\n            out[\"ce_logits\"] = self.decoder(z, tgt_tokens, memory_key_padding_mask=enc_mask)\n        return out\n\n    @torch.no_grad()\n    def inference(self, neural, max_len=100):\n        z, _, enc_mask = self.encoder(neural)\n        return self.decoder.generate(z, max_len=max_len, mem_mask=enc_mask)\n\n    @torch.no_grad()\n    def ctc_decode(self, neural):\n        _, ctc_logits, _ = self.encoder(neural)\n        lp = F.log_softmax(ctc_logits, dim=-1)\n        return [TOKENIZER.ctc_greedy_decode(lp[b]) for b in range(neural.size(0))]\n\n\n# ── Quick smoke test ──────────────────────────────────────────────────────\n_cfg2 = Config(n_encoder_layers=2, n_decoder_layers=2)\n_m    = DSDNLAv2(_cfg2, TOKENIZER.vocab_size)\n_x    = torch.randn(2, 50, 512)\n_t    = torch.randint(2, TOKENIZER.vocab_size, (2, 10))\n_o    = _m(_x, _t)\nprint(\"CTC logits :\", _o[\"ctc_logits\"].shape)\nprint(\"CE  logits :\", _o[\"ce_logits\"].shape)\ndel _cfg2, _m, _x, _t, _o\nprint(\"\\nFull model ✓\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:43.290376Z","iopub.execute_input":"2026-03-03T18:42:43.29106Z","iopub.status.idle":"2026-03-03T18:42:43.79655Z","shell.execute_reply.started":"2026-03-03T18:42:43.291032Z","shell.execute_reply":"2026-03-03T18:42:43.795822Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 6: Dataset & DataLoader","metadata":{}},{"cell_type":"code","source":"class BrainToTextDataset(Dataset):\n    def __init__(self, hdf5_paths, tokenizer, mode=\"train\"):\n        self.tokenizer = tokenizer\n        self.augment   = (mode == \"train\")\n        self._keys: List[Tuple[int, str]] = []\n        self._files: Dict[int, h5py.File] = {}\n        for i, p in enumerate(tqdm(hdf5_paths, desc=f\"Open {mode}\")):\n            f = h5py.File(p, \"r\")\n            self._files[i] = f\n            for k in f.keys():\n                self._keys.append((i, k))\n        print(f\"  → {len(self._keys)} {mode} trials\")\n\n    def __len__(self):  return len(self._keys)\n\n    def __getitem__(self, idx):\n        fi, key  = self._keys[idx]\n        trial    = self._files[fi][key]\n        neural   = torch.tensor(trial[\"input_features\"][:], dtype=torch.float32)\n        text     = trial.attrs.get(\"sentence_label\", \"\")\n        tokens   = self.tokenizer.encode(text)\n        ctc_tgt  = self.tokenizer.encode_ctc(text)\n        if self.augment:\n            neural = self._aug(neural)\n        return {\"neural\": neural, \"tokens\": tokens, \"ctc_tgt\": ctc_tgt, \"text\": text}\n\n    def _aug(self, x):\n        T, C = x.shape\n        if torch.rand(1) < 0.6:   # temporal masking\n            for _ in range(torch.randint(1, 3, (1,)).item()):\n                ml = torch.randint(5, 20, (1,)).item()\n                st = torch.randint(0, max(1, T - ml), (1,)).item()\n                x[st:st+ml] = 0.0\n        if torch.rand(1) < 0.3:   # electrode dropout\n            x = x * torch.bernoulli(torch.full((C,), 0.9))\n        if torch.rand(1) < 0.4:   # gaussian noise\n            x = x + torch.randn_like(x) * 0.04\n        if torch.rand(1) < 0.3:   # time-shift\n            x = torch.roll(x, shifts=torch.randint(-5, 6, (1,)).item(), dims=0)\n        return x\n\n\ndef collate_fn(batch):\n    neurals = [b[\"neural\"]  for b in batch]\n    tokens  = [b[\"tokens\"]  for b in batch]\n    ctcs    = [b[\"ctc_tgt\"] for b in batch]\n    texts   = [b[\"text\"]    for b in batch]\n    B, lengths = len(neurals), [n.size(0) for n in neurals]\n    T_max   = max(lengths)\n    neural_padded = pad_sequence(neurals, batch_first=True, padding_value=0.0)\n    neural_mask   = torch.ones(B, T_max, dtype=torch.bool)\n    for i, l in enumerate(lengths): neural_mask[i, :l] = False\n    return {\n        \"neural\":      neural_padded,\n        \"neural_mask\": neural_mask,\n        \"tokens\":      pad_sequence(tokens, batch_first=True, padding_value=CharTokenizer.PAD_ID),\n        \"ctc_tgt\":     pad_sequence(ctcs,   batch_first=True, padding_value=0),\n        \"ctc_lengths\": torch.tensor([c.size(0) for c in ctcs], dtype=torch.long),\n        \"texts\":       texts,\n    }\n\nprint(\"Dataset ✓\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:43.7979Z","iopub.execute_input":"2026-03-03T18:42:43.79824Z","iopub.status.idle":"2026-03-03T18:42:43.810655Z","shell.execute_reply.started":"2026-03-03T18:42:43.798217Z","shell.execute_reply":"2026-03-03T18:42:43.809922Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 7: Joint CTC + Label-Smoothed CE Loss","metadata":{}},{"cell_type":"code","source":"class JointCTCCELoss(nn.Module):\n    def __init__(self, vocab_size, pad_id, blank_id, ctc_weight=0.3, label_smoothing=0.1):\n        super().__init__()\n        self.ctc_weight = ctc_weight\n        self.ctc = nn.CTCLoss(blank=blank_id, reduction=\"mean\", zero_infinity=True)\n        self.ce  = nn.CrossEntropyLoss(ignore_index=pad_id, label_smoothing=label_smoothing)\n\n    def forward(self, ctc_logits, ce_logits, ctc_tgt, ctc_lengths, tokens, enc_mask=None):\n        B, L, V = ctc_logits.shape\n        lp = F.log_softmax(ctc_logits, dim=-1).permute(1, 0, 2)  # (L, B, V)\n        inp_len = ((~enc_mask).sum(1).long() if enc_mask is not None\n                   else torch.full((B,), L, dtype=torch.long, device=ctc_logits.device))\n        ctc_loss = self.ctc(lp, ctc_tgt, inp_len, ctc_lengths)\n\n        tgt = tokens[:, 1:]\n        T   = min(ce_logits.size(1), tgt.size(1))\n        ce_loss = self.ce(ce_logits[:, :T].reshape(-1, V), tgt[:, :T].reshape(-1))\n\n        total = self.ctc_weight * ctc_loss + (1 - self.ctc_weight) * ce_loss\n        return {\"total\": total, \"ctc\": ctc_loss.detach(), \"ce\": ce_loss.detach()}\n\nprint(\"Loss ✓\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:48.35207Z","iopub.execute_input":"2026-03-03T18:42:48.354477Z","iopub.status.idle":"2026-03-03T18:42:48.361532Z","shell.execute_reply.started":"2026-03-03T18:42:48.354446Z","shell.execute_reply":"2026-03-03T18:42:48.360936Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 8: Trainer","metadata":{}},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, model, cfg):\n        self.model   = model.to(cfg.device)\n        self.cfg     = cfg\n        self.device  = cfg.device\n        self.criterion = JointCTCCELoss(\n            TOKENIZER.vocab_size, CharTokenizer.PAD_ID,\n            CharTokenizer.BLANK_ID, cfg.ctc_weight, cfg.label_smoothing)\n        self.opt = optim.AdamW(model.parameters(), lr=cfg.learning_rate,\n                               weight_decay=cfg.weight_decay, betas=(0.9, 0.98))\n        self.scaler = (torch.amp.GradScaler(\"cuda\")\n                       if cfg.use_amp and cfg.device == \"cuda\" else None)\n        self.scheduler     = None\n        self.best_val_loss = float(\"inf\")\n        self.patience_ctr  = 0\n        self.global_step   = 0\n\n    def _build_sched(self, steps):\n        self.scheduler = optim.lr_scheduler.OneCycleLR(\n            self.opt, max_lr=self.cfg.learning_rate,\n            total_steps=steps * self.cfg.num_epochs,\n            pct_start=self.cfg.warmup_frac, anneal_strategy=\"cos\")\n\n    def _step(self, batch):\n        neural = batch[\"neural\"].to(self.device)\n        nmask  = batch[\"neural_mask\"].to(self.device)\n        tokens = batch[\"tokens\"].to(self.device)\n        ctc_t  = batch[\"ctc_tgt\"].to(self.device)\n        ctc_l  = batch[\"ctc_lengths\"].to(self.device)\n        use_amp = self.scaler is not None\n        with torch.amp.autocast(\"cuda\", enabled=use_amp):\n            out  = self.model(neural, tgt_tokens=tokens[:, :-1], neural_mask=nmask)\n            loss = self.criterion(out[\"ctc_logits\"], out[\"ce_logits\"],\n                                  ctc_t, ctc_l, tokens, out[\"enc_mask\"])\n        self.opt.zero_grad(set_to_none=True)\n        if self.scaler:\n            self.scaler.scale(loss[\"total\"]).backward()\n            self.scaler.unscale_(self.opt)\n            nn.utils.clip_grad_norm_(self.model.parameters(), self.cfg.gradient_clip)\n            self.scaler.step(self.opt); self.scaler.update()\n        else:\n            loss[\"total\"].backward()\n            nn.utils.clip_grad_norm_(self.model.parameters(), self.cfg.gradient_clip)\n            self.opt.step()\n        if self.scheduler and self.global_step > 0: self.scheduler.step()\n        self.global_step += 1\n        return loss\n\n    def train_epoch(self, loader, epoch):\n        self.model.train()\n        acc = {\"total\": 0., \"ctc\": 0., \"ce\": 0.}\n        pbar = tqdm(loader, desc=f\"Epoch {epoch:03d}\")\n        for b in pbar:\n            ls = self._step(b)\n            for k in acc: acc[k] += ls[k].item()\n            pbar.set_postfix({\"loss\": f\"{ls['total'].item():.4f}\"})\n        n = len(loader)\n        return {k: v/n for k, v in acc.items()}\n\n    @torch.no_grad()\n    def val_epoch(self, loader):\n        self.model.eval(); total = 0.\n        for b in tqdm(loader, desc=\"Val\"):\n            neural = b[\"neural\"].to(self.device); nmask = b[\"neural_mask\"].to(self.device)\n            tokens = b[\"tokens\"].to(self.device);  ctc_t = b[\"ctc_tgt\"].to(self.device)\n            ctc_l  = b[\"ctc_lengths\"].to(self.device)\n            out  = self.model(neural, tgt_tokens=tokens[:, :-1], neural_mask=nmask)\n            loss = self.criterion(out[\"ctc_logits\"], out[\"ce_logits\"],\n                                  ctc_t, ctc_l, tokens, out[\"enc_mask\"])\n            total += loss[\"total\"].item()\n        return total / len(loader)\n\n    def save(self, epoch):\n        Path(self.cfg.checkpoint_path).parent.mkdir(parents=True, exist_ok=True)\n        torch.save({\"epoch\": epoch, \"global_step\": self.global_step,\n                    \"model_state_dict\": self.model.state_dict(),\n                    \"optimizer_state_dict\": self.opt.state_dict(),\n                    \"best_val_loss\": self.best_val_loss,\n                    \"config\": CFG.__dict__}, self.cfg.checkpoint_path)\n        print(f\"  ✓ checkpoint → {self.cfg.checkpoint_path}\")\n\n    def fit(self, train_loader, val_loader):\n        self._build_sched(len(train_loader))\n        print(\"\\n\" + \"=\"*60 + \"\\n  DSD-NLA v2 Training\\n\" + \"=\"*60)\n        history = {\"train\": [], \"val\": []}\n        for epoch in range(self.cfg.num_epochs):\n            tr = self.train_epoch(train_loader, epoch)\n            vl = self.val_epoch(val_loader)\n            history[\"train\"].append(tr[\"total\"])\n            history[\"val\"].append(vl)\n            print(f\"Epoch {epoch:03d}  train={tr['total']:.4f} \"\n                  f\"(ctc={tr['ctc']:.4f} ce={tr['ce']:.4f})  val={vl:.4f}\")\n            if vl < self.best_val_loss:\n                self.best_val_loss = vl; self.patience_ctr = 0; self.save(epoch)\n            else:\n                self.patience_ctr += 1\n                print(f\"  No improvement ({self.patience_ctr}/{self.cfg.patience})\")\n                if self.patience_ctr >= self.cfg.patience:\n                    print(f\"  Early stopping at epoch {epoch}.\"); break\n        print(f\"\\nBest val loss: {self.best_val_loss:.4f}\")\n        return history\n\nprint(\"Trainer ✓\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:50.416846Z","iopub.execute_input":"2026-03-03T18:42:50.41715Z","iopub.status.idle":"2026-03-03T18:42:50.433808Z","shell.execute_reply.started":"2026-03-03T18:42:50.417123Z","shell.execute_reply":"2026-03-03T18:42:50.433134Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 9: Build Datasets & Train","metadata":{}},{"cell_type":"code","source":"train_files = sorted(glob.glob(f\"{CFG.data_dir}/t15.*/data_train.hdf5\"))\nval_files   = sorted(glob.glob(f\"{CFG.data_dir}/t15.*/data_val.hdf5\"))\nprint(f\"Train files: {len(train_files)}  |  Val files: {len(val_files)}\")\n\ntrain_ds = BrainToTextDataset(train_files, TOKENIZER, mode=\"train\")\nval_ds   = BrainToTextDataset(val_files,   TOKENIZER, mode=\"val\")\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, shuffle=True,\n                          collate_fn=collate_fn, num_workers=2, pin_memory=True)\nval_loader   = DataLoader(val_ds,   batch_size=CFG.batch_size, shuffle=False,\n                          collate_fn=collate_fn, num_workers=2, pin_memory=True)\n\nmodel = DSDNLAv2(CFG, TOKENIZER.vocab_size)\nprint(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n\nRUN_TRAINING = True\nif RUN_TRAINING:\n    trainer = Trainer(model, CFG)\n    history = trainer.fit(train_loader, val_loader)\n\nprint(\"\\nFiles:\", os.listdir(\"/kaggle/working\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:51.040472Z","iopub.execute_input":"2026-03-03T18:42:51.040775Z","iopub.status.idle":"2026-03-03T18:42:51.061059Z","shell.execute_reply.started":"2026-03-03T18:42:51.04075Z","shell.execute_reply":"2026-03-03T18:42:51.060071Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 10: Training Curves","metadata":{}},{"cell_type":"code","source":"if RUN_TRAINING and \"history\" in dir():\n    fig, ax = plt.subplots(figsize=(10, 4))\n    ax.plot(history[\"train\"], label=\"Train\", lw=2)\n    ax.plot(history[\"val\"],   label=\"Val\",   lw=2, ls=\"--\")\n    ax.set_xlabel(\"Epoch\"); ax.set_ylabel(\"Loss\")\n    ax.set_title(\"DSD-NLA v2 — Training Curves\")\n    ax.legend(); ax.grid(alpha=0.3); plt.tight_layout()\n    plt.savefig(\"/kaggle/working/training_curves.png\", dpi=150)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:51.426927Z","iopub.execute_input":"2026-03-03T18:42:51.427221Z","iopub.status.idle":"2026-03-03T18:42:51.434421Z","shell.execute_reply.started":"2026-03-03T18:42:51.427196Z","shell.execute_reply":"2026-03-03T18:42:51.433494Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 11: Post-Processing","metadata":{}},{"cell_type":"code","source":"def normalize_for_eval(text: str) -> str:\n    text = text.lower().replace(\"\\u2019\", \"'\")\n    text = re.sub(r\"[^a-z0-9' ]\", \" \", text)\n    return re.sub(r\"\\s+\", \" \", text).strip()\n\ndef clean_text(text: str) -> str:\n    text = text.lower()\n    text = re.sub(r\"([aeiou])\\1{2,}\", r\"\\1\\1\", text)\n    return re.sub(r\"\\s+\", \" \", text).strip()\n\nprint(normalize_for_eval(\"Hello, World! It's great.\"))\nprint(clean_text(\"soooo goood\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:51.778466Z","iopub.execute_input":"2026-03-03T18:42:51.77908Z","iopub.status.idle":"2026-03-03T18:42:51.784247Z","shell.execute_reply.started":"2026-03-03T18:42:51.779051Z","shell.execute_reply":"2026-03-03T18:42:51.783517Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 12: Inference & Submission","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef generate_submission(checkpoint_path, test_data_dir, output_path,\n                         device=\"cuda\", use_ctc=True, max_len=100):\n    ckpt = torch.load(checkpoint_path, map_location=device)\n    cfg_d = ckpt.get(\"config\", {})\n    inf_cfg = Config(\n        n_channels=cfg_d.get(\"n_channels\", CFG.n_channels),\n        d_model=cfg_d.get(\"d_model\", CFG.d_model),\n        n_encoder_layers=cfg_d.get(\"n_encoder_layers\", CFG.n_encoder_layers),\n        n_decoder_layers=cfg_d.get(\"n_decoder_layers\", CFG.n_decoder_layers),\n        n_heads=cfg_d.get(\"n_heads\", CFG.n_heads),\n        conv_kernel=cfg_d.get(\"conv_kernel\", CFG.conv_kernel),\n        ff_expansion=cfg_d.get(\"ff_expansion\", CFG.ff_expansion),\n        dropout=0.0, stochastic_depth_prob=0.0)\n    m = DSDNLAv2(inf_cfg, TOKENIZER.vocab_size)\n    m.load_state_dict(ckpt[\"model_state_dict\"])\n    m = m.to(device).eval()\n    print(f\"Model loaded ({sum(p.numel() for p in m.parameters()):,} params)\")\n\n    test_files = sorted(glob.glob(f\"{test_data_dir}/t15.*/data_test.hdf5\"))\n    print(f\"Test files: {len(test_files)}\")\n    preds = []\n    for fp in tqdm(test_files, desc=\"Files\"):\n        with h5py.File(fp, \"r\") as f:\n            for key in tqdm(sorted(f.keys()), desc=Path(fp).parent.name, leave=False):\n                neural = torch.tensor(f[key][\"input_features\"][:],\n                                      dtype=torch.float32).unsqueeze(0).to(device)\n                if use_ctc:\n                    raw = m.ctc_decode(neural)[0]\n                else:\n                    tok = m.inference(neural, max_len=max_len)\n                    raw = TOKENIZER.decode(tok[0])\n                preds.append(normalize_for_eval(clean_text(raw)))\n\n    df = pd.DataFrame({\"id\": range(len(preds)), \"text\": preds})\n    df.to_csv(output_path, index=False)\n    print(f\"\\n{len(preds)} predictions → {output_path}\")\n    print(df.head(10).to_string(index=False))\n    return df\n\nsubmission_df = generate_submission(\n    CFG.checkpoint_path, CFG.data_dir, CFG.submission_path,\n    device=CFG.device, use_ctc=CFG.ctc_decode)\nprint(\"\\n✅ submission.csv ready!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:52.238729Z","iopub.execute_input":"2026-03-03T18:42:52.23957Z","iopub.status.idle":"2026-03-03T18:42:52.260541Z","shell.execute_reply.started":"2026-03-03T18:42:52.239539Z","shell.execute_reply":"2026-03-03T18:42:52.259623Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 13: Validation WER Evaluation","metadata":{}},{"cell_type":"code","source":"def simple_wer(hyp, ref):\n    h, r = hyp.split(), ref.split()\n    if not r: return 1.0 if h else 0.0\n    d = np.zeros((len(h)+1, len(r)+1), dtype=int)\n    for i in range(len(h)+1): d[i][0] = i\n    for j in range(len(r)+1): d[0][j] = j\n    for i in range(1, len(h)+1):\n        for j in range(1, len(r)+1):\n            c = 0 if h[i-1] == r[j-1] else 1\n            d[i][j] = min(d[i-1][j]+1, d[i][j-1]+1, d[i-1][j-1]+c)\n    return d[len(h)][len(r)] / max(len(r), 1)\n\n@torch.no_grad()\ndef evaluate_val(n=20):\n    if not Path(CFG.checkpoint_path).exists():\n        print(\"No checkpoint found.\"); return\n    ckpt = torch.load(CFG.checkpoint_path, map_location=CFG.device)\n    em = DSDNLAv2(CFG, TOKENIZER.vocab_size)\n    em.load_state_dict(ckpt[\"model_state_dict\"])\n    em = em.to(CFG.device).eval()\n    vf  = sorted(glob.glob(f\"{CFG.data_dir}/t15.*/data_val.hdf5\"))\n    rows, cnt = [], 0\n    for fp in vf:\n        if cnt >= n: break\n        with h5py.File(fp, \"r\") as f:\n            for key in sorted(f.keys()):\n                if cnt >= n: break\n                nr  = torch.tensor(f[key][\"input_features\"][:],\n                                   dtype=torch.float32).unsqueeze(0).to(CFG.device)\n                gt  = normalize_for_eval(f[key].attrs.get(\"sentence_label\", \"\"))\n                pc  = normalize_for_eval(clean_text(em.ctc_decode(nr)[0]))\n                tok = em.inference(nr, max_len=100)\n                pa  = normalize_for_eval(clean_text(TOKENIZER.decode(tok[0])))\n                rows.append({\"GT\": gt, \"CTC\": pc, \"ATT\": pa,\n                             \"WER_CTC\": simple_wer(pc, gt),\n                             \"WER_ATT\": simple_wer(pa, gt)})\n                cnt += 1\n    df = pd.DataFrame(rows)\n    print(f\"\\n=== Val samples (n={n}) ===\")\n    for _, r in df.iterrows():\n        print(f\"  GT : {r['GT']}\")\n        print(f\"  CTC: {r['CTC']}  WER={r['WER_CTC']:.2f}\")\n        print(f\"  ATT: {r['ATT']}  WER={r['WER_ATT']:.2f}\\n\")\n    print(f\"Mean WER — CTC: {df['WER_CTC'].mean():.4f}\")\n    print(f\"Mean WER — ATT: {df['WER_ATT'].mean():.4f}\")\n    return df\n\neval_df = evaluate_val(20)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:52.662738Z","iopub.execute_input":"2026-03-03T18:42:52.663034Z","iopub.status.idle":"2026-03-03T18:42:52.674168Z","shell.execute_reply.started":"2026-03-03T18:42:52.663008Z","shell.execute_reply":"2026-03-03T18:42:52.67364Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 14: Ablation & Future Work\n\n| Improvement | Expected ΔWER | Status |\n|---|---|---|\n| CTC loss (v1→v2) | ↓ 5–10% | ✅ implemented |\n| Conformer encoder | ↓ 3–7% | ✅ implemented |\n| Stochastic depth | ↓ 1–2% | ✅ implemented |\n| Label smoothing | ↓ 1% | ✅ implemented |\n| CTC beam search + n-gram LM | ↓ 5–8% | 🔧 next step |\n| Cross-session normalisation | ↓ 2–5% | 🔧 next step |\n| Scale to d_model=768 | ↓ 3–7% | 🔧 next step |\n| Model ensemble (5× seeds) | ↓ 3–5% | 🔧 next step |\n\n### Recommended Next Steps\n1. **CTC beam search with kenlm**: Install `ctcdecode`/`flashlight-text` and build a 4-gram LM from training transcripts.\n2. **Cross-session z-score**: Compute μ/σ per channel per session — reduces distribution shift.\n3. **Phoneme auxiliary head**: Add a phoneme CTC head as an auxiliary loss (diphone approach from BrainBench '24 leaders).\n4. **Larger model**: Scale `d_model=768, n_encoder=12` once GPU budget allows.\n5. **Ensemble**: Average CTC logits across 3–5 independently trained seeds.","metadata":{}},{"cell_type":"markdown","source":"## Cell 15: Final Output Check","metadata":{}},{"cell_type":"code","source":"print(\"Files in /kaggle/working:\")\nfor f in sorted(os.listdir(\"/kaggle/working\")):\n    sz = os.path.getsize(f\"/kaggle/working/{f}\")\n    print(f\"  {f:<40} {sz/1e6:.2f} MB\")\nassert \"submission.csv\" in os.listdir(\"/kaggle/working\"), \"MISSING submission.csv!\"\nprint(\"\\n✅ submission.csv present — ready for Kaggle submission!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:42:53.231449Z","iopub.execute_input":"2026-03-03T18:42:53.231721Z","iopub.status.idle":"2026-03-03T18:42:53.238782Z","shell.execute_reply.started":"2026-03-03T18:42:53.2317Z","shell.execute_reply":"2026-03-03T18:42:53.237842Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}