{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🧠 Brain-to-Text Decoding — MSc Thesis Implementation\n**Student:** Uzma Rehman  \n**Dataset:** Brain-to-Text Benchmark '25 (T15, 256-channel Utah intracortical array)  \n**Method:** SSL-Pretrained Transformer Encoder + BiGRU CTC Decoder  \n**Paper:** BIT Framework — arxiv 2511.21740  \n\n---\n\n## System Overview\n```\nNeural Data (HDF5)\n      ↓\nPreprocessing (z-score + time-patch windowing)\n      ↓\nSSL-Pretrained Transformer Encoder (6 layers, 512 dim)\n      ↓\nBiGRU CTC Decoder (3 layers, 512 hidden)\n      ↓\nDecoded Text → WER / CER / BLEU evaluation\n```\n\n**All model weights are loaded from the `brain-to-text-checkpoints` dataset — no training is performed in this notebook.**","metadata":{}},{"cell_type":"markdown","source":"## Section 1 — Imports & Setup","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport numpy as np\nimport h5py\nimport glob\nimport os\nimport json\nimport math\nimport string\nimport time\nimport csv\nfrom torch.utils.data import Dataset, DataLoader\nimport matplotlib.pyplot as plt\n\n# Device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"✅ Libraries imported\")\nprint(f\"✅ Device: {'GPU — ' + torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'CPU'}\")\n\n# Checkpoint dataset path\nCKPT_DIR = '/kaggle/input/datasets/uzmarehman/brain-to-text-checkpoints'\nWORK_DIR  = '/kaggle/working'\nos.makedirs(WORK_DIR, exist_ok=True)\n\n# Verify dataset is attached\nif os.path.exists(CKPT_DIR):\n    files = os.listdir(CKPT_DIR)\n    print(f\"✅ Checkpoint dataset found: {len(files)} files\")\nelse:\n    print(\"❌ brain-to-text-checkpoints dataset not attached — add it via + Add Input\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:17:57.889448Z","iopub.execute_input":"2026-06-27T08:17:57.891194Z","iopub.status.idle":"2026-06-27T08:17:57.904085Z","shell.execute_reply.started":"2026-06-27T08:17:57.891079Z","shell.execute_reply":"2026-06-27T08:17:57.903474Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 2 — Configuration","metadata":{}},{"cell_type":"code","source":"# Load config from saved dataset\nwith open(f'{CKPT_DIR}/config.json') as f:\n    CONFIG = json.load(f)\n\nprint(\"✅ Config loaded:\")\nfor k, v in CONFIG.items():\n    if k != 'dataset_stats':\n        print(f\"  {k:20s}: {v}\")\nprint(f\"\\n  dataset_stats:\")\nfor k, v in CONFIG.get('dataset_stats', {}).items():\n    print(f\"    {k:20s}: {v}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:17:57.905718Z","iopub.execute_input":"2026-06-27T08:17:57.906040Z","iopub.status.idle":"2026-06-27T08:17:57.913473Z","shell.execute_reply.started":"2026-06-27T08:17:57.906019Z","shell.execute_reply":"2026-06-27T08:17:57.912892Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 3 — Dataset & Preprocessing\n\n### 3.1 Helper Functions\n\nThe preprocessing pipeline converts raw HDF5 neural data into batched tensors:\n1. **Z-score normalisation** — rescale each channel to mean=0, std=1\n2. **Time-patch windowing** — group 4×20ms timesteps into one patch token (80ms per patch)\n3. **Lazy loading** — load trials on-demand to avoid RAM overflow\n4. **Padding + collate** — pad variable-length sequences to batch uniformly","metadata":{}},{"cell_type":"code","source":"def make_patches(feature_array, patch_size=4):\n    T, C   = feature_array.shape\n    T_trim = (T // patch_size) * patch_size\n    return feature_array[:T_trim].reshape(T_trim // patch_size, patch_size * C)\n\ndef collate_fn(batch):\n    features, phonemes, transcriptions = zip(*batch)\n    max_len  = max(f.shape[0] for f in features)\n    feat_dim = features[0].shape[1]\n    padded   = torch.zeros(len(features), max_len, feat_dim)\n    lengths  = []\n    for i, f in enumerate(features):\n        padded[i, :f.shape[0], :] = f\n        lengths.append(f.shape[0])\n    return padded, torch.stack(phonemes), torch.stack(transcriptions), torch.tensor(lengths)\n\nprint(\"✅ Helper functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:17:57.914232Z","iopub.execute_input":"2026-06-27T08:17:57.914454Z","iopub.status.idle":"2026-06-27T08:17:57.920477Z","shell.execute_reply.started":"2026-06-27T08:17:57.914435Z","shell.execute_reply":"2026-06-27T08:17:57.919949Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3.2 Dataset Class","metadata":{}},{"cell_type":"code","source":"class BrainToTextDataset(Dataset):\n    \"\"\"\n    Memory-efficient lazy-loading dataset for Brain-to-Text '25.\n    Only file paths are stored at init time; data is loaded on-demand.\n    Handles missing labels in competition test files gracefully.\n    \"\"\"\n    def __init__(self, file_paths, patch_size=4, mean=None, std=None):\n        self.patch_size = patch_size\n        self.index      = []\n        for path in file_paths:\n            with h5py.File(path, 'r') as f:\n                for key in sorted(f.keys()):\n                    self.index.append((path, key))\n        if mean is None or std is None:\n            sample_data = []\n            for path in file_paths[:5]:\n                with h5py.File(path, 'r') as f:\n                    for key in sorted(f.keys()):\n                        sample_data.append(f[key]['input_features'][:])\n            all_data   = np.concatenate(sample_data, axis=0)\n            self.mean  = all_data.mean(axis=0, keepdims=True)\n            self.std   = all_data.std(axis=0,  keepdims=True) + 1e-8\n            del sample_data, all_data\n        else:\n            self.mean = mean\n            self.std  = std\n\n    def __len__(self):\n        return len(self.index)\n\n    def __getitem__(self, idx):\n        path, key = self.index[idx]\n        with h5py.File(path, 'r') as f:\n            feat  = f[key]['input_features'][:]\n            phon  = f[key]['seq_class_ids'][:] \\\n                    if 'seq_class_ids' in f[key] \\\n                    else np.zeros(500, dtype=np.int64)\n            trans = f[key]['transcription'][:] \\\n                    if 'transcription'  in f[key] \\\n                    else np.zeros(500, dtype=np.int64)\n        feat = (feat - self.mean) / self.std\n        feat = make_patches(feat, self.patch_size)\n        return (torch.tensor(feat,  dtype=torch.float32),\n                torch.tensor(phon,  dtype=torch.long),\n                torch.tensor(trans, dtype=torch.long))\n\nprint(\"✅ Dataset class defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:17:57.921249Z","iopub.execute_input":"2026-06-27T08:17:57.921549Z","iopub.status.idle":"2026-06-27T08:17:57.930318Z","shell.execute_reply.started":"2026-06-27T08:17:57.921528Z","shell.execute_reply":"2026-06-27T08:17:57.929497Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3.3 Load Data","metadata":{}},{"cell_type":"code","source":"BASE        = '/kaggle/input/competitions/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final'\ntrain_files = sorted(glob.glob(f'{BASE}/*/data_train.hdf5'))\nval_files   = sorted(glob.glob(f'{BASE}/*/data_val.hdf5'))\ntest_files  = sorted(glob.glob(f'{BASE}/*/data_test.hdf5'))\n\n# Load normalisation stats from permanent dataset\nmean_loaded = np.load(f'{CKPT_DIR}/norm_mean.npy')\nstd_loaded  = np.load(f'{CKPT_DIR}/norm_std.npy')\nprint(f\"✅ Norm stats loaded — shape: {mean_loaded.shape}\")\n\ntrain_dataset = BrainToTextDataset(train_files, CONFIG['patch_size'], mean_loaded, std_loaded)\nval_dataset   = BrainToTextDataset(val_files,   CONFIG['patch_size'], mean_loaded, std_loaded)\ntest_dataset  = BrainToTextDataset(test_files,  CONFIG['patch_size'], mean_loaded, std_loaded)\n\ntrain_loader = DataLoader(train_dataset, batch_size=CONFIG['batch_size'], shuffle=True,  collate_fn=collate_fn)\nval_loader   = DataLoader(val_dataset,   batch_size=CONFIG['batch_size'], shuffle=False, collate_fn=collate_fn)\ntest_loader  = DataLoader(test_dataset,  batch_size=CONFIG['batch_size'], shuffle=False, collate_fn=collate_fn)\n\nprint(f\"✅ Datasets loaded:\")\nprint(f\"   Train: {len(train_dataset):,} trials | Val: {len(val_dataset):,} | Test: {len(test_dataset):,}\")\nprint(f\"   Sessions: {len(train_files)} train | {len(val_files)} val | {len(test_files)} test\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:17:57.931990Z","iopub.execute_input":"2026-06-27T08:17:57.932576Z","iopub.status.idle":"2026-06-27T08:18:06.861804Z","shell.execute_reply.started":"2026-06-27T08:17:57.932552Z","shell.execute_reply":"2026-06-27T08:18:06.861075Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 4 — Model Architecture\n\n### 4.1 SSL Transformer Encoder\n\nThe encoder consists of:\n- **SubjectReadIn**: Linear(2048→512) + LayerNorm — maps subject-specific patch features to model dimension\n- **PositionalEncoding**: Sinusoidal position fingerprints — gives the transformer a sense of order\n- **6× TransformerBlock**: Multi-head attention (8 heads) + FFN (GELU) + residual connections + LayerNorm\n- **Final LayerNorm**: Stabilises output representations","metadata":{}},{"cell_type":"code","source":"class SubjectReadIn(nn.Module):\n    def __init__(self, patch_dim=2048, model_dim=512):\n        super().__init__()\n        self.linear = nn.Linear(patch_dim, model_dim)\n        self.norm   = nn.LayerNorm(model_dim)\n    def forward(self, x):\n        return self.norm(self.linear(x))\n\nclass SubjectReadOut(nn.Module):\n    def __init__(self, model_dim=512, patch_dim=2048):\n        super().__init__()\n        self.linear = nn.Linear(model_dim, patch_dim)\n    def forward(self, x):\n        return self.linear(x)\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, model_dim=512, max_len=1000, dropout=0.1):\n        super().__init__()\n        self.dropout = nn.Dropout(dropout)\n        pe           = torch.zeros(max_len, model_dim)\n        position     = torch.arange(0, max_len).unsqueeze(1).float()\n        div_term     = torch.exp(torch.arange(0, model_dim, 2).float() *\n                       (-math.log(10000.0) / model_dim))\n        pe[:, 0::2]  = torch.sin(position * div_term)\n        pe[:, 1::2]  = torch.cos(position * div_term)\n        self.register_buffer('pe', pe.unsqueeze(0))\n    def forward(self, x):\n        return self.dropout(x + self.pe[:, :x.shape[1], :])\n\nclass TransformerBlock(nn.Module):\n    def __init__(self, model_dim=512, num_heads=8, ffn_dim=2048, dropout=0.1):\n        super().__init__()\n        self.attention = nn.MultiheadAttention(model_dim, num_heads,\n                                               dropout=dropout, batch_first=True)\n        self.ffn = nn.Sequential(\n            nn.Linear(model_dim, ffn_dim), nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(ffn_dim, model_dim), nn.Dropout(dropout))\n        self.norm1 = nn.LayerNorm(model_dim)\n        self.norm2 = nn.LayerNorm(model_dim)\n    def forward(self, x, key_padding_mask=None):\n        attn_out, _ = self.attention(x, x, x, key_padding_mask=key_padding_mask)\n        x = self.norm1(x + attn_out)\n        return self.norm2(x + self.ffn(x))\n\nclass NeuralTransformerEncoder(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        self.read_in      = SubjectReadIn(config['patch_dim'], config['model_dim'])\n        self.pos_encoding = PositionalEncoding(config['model_dim'], dropout=config['dropout'])\n        self.blocks       = nn.ModuleList([\n            TransformerBlock(config['model_dim'], config['num_heads'],\n                             config['ffn_dim'],   config['dropout'])\n            for _ in range(config['num_layers'])\n        ])\n        self.final_norm   = nn.LayerNorm(config['model_dim'])\n    def forward(self, x, lengths=None):\n        key_padding_mask = None\n        if lengths is not None:\n            B, T, _          = x.shape\n            key_padding_mask = torch.arange(T, device=x.device)\\\n                .unsqueeze(0) >= lengths.unsqueeze(1)\n        x = self.read_in(x)\n        x = self.pos_encoding(x)\n        for block in self.blocks:\n            x = block(x, key_padding_mask)\n        return self.final_norm(x)\n\nclass PatchMasking(nn.Module):\n    def __init__(self, mask_ratio=0.75):\n        super().__init__()\n        self.mask_ratio = mask_ratio\n    def forward(self, x, lengths=None):\n        B, T, D = x.shape\n        results = []\n        for i in range(B):\n            real_len = lengths[i].item() if lengths is not None else T\n            num_mask = int(real_len * self.mask_ratio)\n            perm     = torch.randperm(real_len, device=x.device)\n            mask     = torch.zeros(T, dtype=torch.bool, device=x.device)\n            mask[perm[:num_mask]] = True\n            results.append(mask)\n        masks         = torch.stack(results)\n        x_masked      = x.clone()\n        x_masked[masks] = 0.0\n        return x_masked, masks\n\nclass SSLPretrainingModel(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        self.masking  = PatchMasking(mask_ratio=config['mask_ratio'])\n        self.encoder  = NeuralTransformerEncoder(config)\n        self.read_out = SubjectReadOut(config['model_dim'], config['patch_dim'])\n    def forward(self, x, lengths=None):\n        x_masked, masks = self.masking(x, lengths)\n        encoded         = self.encoder(x_masked, lengths)\n        reconstructed   = self.read_out(encoded)\n        loss            = self.compute_loss(x, reconstructed, masks, lengths)\n        return loss, reconstructed, masks\n    def compute_loss(self, original, reconstructed, masks, lengths):\n        B, T, D = original.shape\n        losses  = []\n        for i in range(B):\n            real_len    = lengths[i].item() if lengths is not None else T\n            real_mask   = masks[i, :real_len]\n            if real_mask.sum() == 0:\n                continue\n            orig_masked = original[i, :real_len][real_mask]\n            rec_masked  = reconstructed[i, :real_len][real_mask]\n            losses.append(((orig_masked - rec_masked) ** 2).mean())\n        return torch.stack(losses).mean()\n\nprint(\"✅ SSL Transformer Encoder classes defined\")\nprint(f\"   Layers: {CONFIG['num_layers']} | Heads: {CONFIG['num_heads']} | Dim: {CONFIG['model_dim']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:18:06.863253Z","iopub.execute_input":"2026-06-27T08:18:06.863534Z","iopub.status.idle":"2026-06-27T08:18:06.881331Z","shell.execute_reply.started":"2026-06-27T08:18:06.863512Z","shell.execute_reply":"2026-06-27T08:18:06.880555Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4.2 GRU+CTC Decoder\n\nThe decoder consists of:\n- **BiGRU** (3 layers, 512 hidden, bidirectional) — reads both forward and backward context\n- **CTC Loss** — enables sequence decoding without explicit alignment\n- **Character vocabulary** (96 tokens + blank) — character-level decoding handles any word","metadata":{}},{"cell_type":"code","source":"CHARS     = [' '] + list(string.ascii_lowercase + string.ascii_uppercase +\n             string.digits + string.punctuation)\nBLANK     = 0\nCHAR2IDX  = {c: i+1 for i, c in enumerate(CHARS)}\nIDX2CHAR  = {i+1: c for i, c in enumerate(CHARS)}\nIDX2CHAR[0] = ''\nVOCAB_SIZE  = len(CHARS) + 1\n\nclass GRUDecoder(nn.Module):\n    def __init__(self, input_dim=512, hidden_dim=512,\n                 vocab_size=VOCAB_SIZE, num_layers=3, dropout=0.2):\n        super().__init__()\n        self.gru     = nn.GRU(input_dim, hidden_dim, num_layers=num_layers,\n                               batch_first=True, bidirectional=True,\n                               dropout=dropout if num_layers > 1 else 0.0)\n        self.dropout = nn.Dropout(dropout)\n        self.fc      = nn.Linear(hidden_dim * 2, vocab_size)\n    def forward(self, x):\n        out, _ = self.gru(x)\n        return self.fc(self.dropout(out))\n\nclass BrainToTextCTC(nn.Module):\n    def __init__(self, encoder, decoder):\n        super().__init__()\n        self.encoder = encoder\n        self.decoder = decoder\n    def forward(self, neural_input, lengths):\n        encoded = self.encoder(neural_input, lengths)\n        logits  = self.decoder(encoded)\n        return logits, lengths\n\nprint(f\"✅ GRU+CTC Decoder classes defined\")\nprint(f\"   Vocab size: {VOCAB_SIZE} | Blank token: {BLANK}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:18:06.882271Z","iopub.execute_input":"2026-06-27T08:18:06.883498Z","iopub.status.idle":"2026-06-27T08:18:06.893096Z","shell.execute_reply.started":"2026-06-27T08:18:06.883475Z","shell.execute_reply":"2026-06-27T08:18:06.890225Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 5 — Load Pretrained Weights","metadata":{}},{"cell_type":"code","source":"# Load SSL pretrained encoder\nssl_model = SSLPretrainingModel(CONFIG).to(device)\nckpt      = torch.load(f'{CKPT_DIR}/ssl_best.pt', map_location=device, weights_only=False)\nstate_dict = {k.replace('module.', ''): v for k, v in ckpt['model_state_dict'].items()}\nssl_model.load_state_dict(state_dict)\nencoder = ssl_model.encoder\nencoder.eval()\nprint(f\"✅ SSL encoder loaded — epoch {ckpt['epoch']}, val_loss={ckpt['val_loss']:.4f}\")\nprint(f\"   Parameters: {sum(p.numel() for p in encoder.parameters()):,}\")\n\n# Build and load main CTC model (best Phase 2)\ngru_decoder = GRUDecoder(input_dim=CONFIG['model_dim'], hidden_dim=512,\n                          vocab_size=VOCAB_SIZE, num_layers=3).to(device)\nctc_model   = BrainToTextCTC(encoder, gru_decoder).to(device)\n\nckpt_ctc   = torch.load(f'{CKPT_DIR}/best_ctc_p2.pt', map_location=device, weights_only=False)\nstate_dict = {k.replace('module.', ''): v for k, v in ckpt_ctc['model_state_dict'].items()}\nctc_model.load_state_dict(state_dict)\nctc_model.eval()\nprint(f\"\\n✅ CTC model loaded — epoch {ckpt_ctc['epoch']}, val_WER={ckpt_ctc['val_wer']*100:.1f}%\")\nprint(f\"   Encoder params: {sum(p.numel() for p in encoder.parameters()):,}\")\nprint(f\"   Decoder params: {sum(p.numel() for p in gru_decoder.parameters()):,}\")\nprint(f\"   Total params:   {sum(p.numel() for p in ctc_model.parameters()):,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:18:06.895582Z","iopub.execute_input":"2026-06-27T08:18:06.897309Z","iopub.status.idle":"2026-06-27T08:18:09.948380Z","shell.execute_reply.started":"2026-06-27T08:18:06.897284Z","shell.execute_reply":"2026-06-27T08:18:09.947483Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 6 — Evaluation Functions","metadata":{}},{"cell_type":"code","source":"def ctc_greedy_decode(log_probs):\n    tokens    = log_probs.argmax(-1).tolist()\n    collapsed = [t for i, t in enumerate(tokens)\n                 if t != BLANK and (i == 0 or t != tokens[i-1])]\n    return ''.join([IDX2CHAR.get(t, '') for t in collapsed]).strip()\n\ndef compute_wer(reference, hypothesis):\n    if isinstance(reference, str):\n        ref_tokens = reference.lower().split()\n        hyp_tokens = hypothesis.lower().split()\n    else:\n        ref_tokens = reference\n        hyp_tokens = hypothesis\n    r, h = len(ref_tokens), len(hyp_tokens)\n    d    = np.zeros((r+1, h+1), dtype=int)\n    for i in range(r+1): d[i][0] = i\n    for j in range(h+1): d[0][j] = j\n    for i in range(1, r+1):\n        for j in range(1, h+1):\n            if ref_tokens[i-1] == hyp_tokens[j-1]: d[i][j] = d[i-1][j-1]\n            else: d[i][j] = 1 + min(d[i-1][j], d[i][j-1], d[i-1][j-1])\n    return d[r][h] / max(len(ref_tokens), 1)\n\ndef compute_bleu(reference, hypothesis, n=1):\n    ref_ngrams = [reference[i:i+n] for i in range(len(reference)-n+1)]\n    hyp_ngrams = [hypothesis[i:i+n] for i in range(len(hypothesis)-n+1)]\n    if not hyp_ngrams: return 0.0\n    return sum(1 for ng in hyp_ngrams if ng in ref_ngrams) / max(len(hyp_ngrams), 1)\n\ndef full_evaluation(model, loader, device, split_name='val'):\n    model.eval()\n    all_refs, all_hyps = [], []\n    is_labelled = split_name in ('val', 'train')\n    print(f\"Evaluating on {split_name} set...\")\n    with torch.no_grad():\n        for batch_idx, (feat, _, trans, lengths) in enumerate(loader):\n            feat    = feat.to(device)\n            lengths = lengths.to(device)\n            logits, _ = model(feat, lengths)\n            log_probs = torch.nn.functional.log_softmax(logits, dim=-1)\n            for i in range(feat.shape[0]):\n                pred_text = ctc_greedy_decode(log_probs[i, :lengths[i]].cpu())\n                all_hyps.append(pred_text)\n                if is_labelled:\n                    ref_chars = [chr(c) for c in trans[i].numpy() if c > 0]\n                    all_refs.append(''.join(ref_chars).strip())\n            if (batch_idx + 1) % 50 == 0:\n                print(f\"  Processed {(batch_idx+1)*8} trials...\")\n    if is_labelled:\n        wer_scores   = [compute_wer(r, h) for r, h in zip(all_refs, all_hyps)]\n        cer_scores   = [compute_wer(list(r), list(h)) for r, h in zip(all_refs, all_hyps)]\n        bleu1_scores = [compute_bleu(r, h, n=1) for r, h in zip(all_refs, all_hyps)]\n        bleu4_scores = [compute_bleu(r, h, n=4) for r, h in zip(all_refs, all_hyps)]\n        results = {\n            'split': split_name, 'n_trials': len(all_refs),\n            'WER':    round(np.mean(wer_scores)   * 100, 2),\n            'CER':    round(np.mean(cer_scores)   * 100, 2),\n            'BLEU-1': round(np.mean(bleu1_scores) * 100, 2),\n            'BLEU-4': round(np.mean(bleu4_scores) * 100, 2),\n        }\n        print(f\"\\n{'='*45}\")\n        print(f\"  RESULTS — {split_name.upper()} SET\")\n        print(f\"{'='*45}\")\n        print(f\"  Trials: {results['n_trials']:,}\")\n        print(f\"  WER:    {results['WER']}%\")\n        print(f\"  CER:    {results['CER']}%\")\n        print(f\"  BLEU-1: {results['BLEU-1']}%\")\n        print(f\"  BLEU-4: {results['BLEU-4']}%\")\n        print(f\"{'='*45}\")\n        print(f\"\\nSample predictions:\")\n        for i in range(min(5, len(all_refs))):\n            print(f\"  REF:  '{all_refs[i]}'\")\n            print(f\"  PRED: '{all_hyps[i]}'\")\n            print(f\"  WER: {wer_scores[i]:.2f} | CER: {cer_scores[i]:.2f}\\n\")\n        return results, all_refs, all_hyps\n    else:\n        results = {'split': split_name, 'n_trials': len(all_hyps),\n                   'WER': 'N/A', 'CER': 'N/A', 'BLEU-1': 'N/A', 'BLEU-4': 'N/A'}\n        print(f\"\\n✅ {len(all_hyps):,} predictions decoded (no labels in test set)\")\n        for i in range(min(3, len(all_hyps))):\n            print(f\"  PRED: '{all_hyps[i]}'\")\n        return results, [], all_hyps\n\nprint(\"✅ Evaluation functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:18:09.949428Z","iopub.execute_input":"2026-06-27T08:18:09.949728Z","iopub.status.idle":"2026-06-27T08:18:09.965559Z","shell.execute_reply.started":"2026-06-27T08:18:09.949706Z","shell.execute_reply":"2026-06-27T08:18:09.964907Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 7 — Main Results\n\nEvaluating the main model (SSL Transformer + BiGRU CTC, mask_ratio=0.75) on the validation and test sets.","metadata":{}},{"cell_type":"code","source":"# Evaluate on validation set (has labels)\nval_results, val_refs, val_hyps = full_evaluation(ctc_model, val_loader, device, 'val')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:18:09.966573Z","iopub.execute_input":"2026-06-27T08:18:09.966836Z","iopub.status.idle":"2026-06-27T08:19:16.289788Z","shell.execute_reply.started":"2026-06-27T08:18:09.966807Z","shell.execute_reply":"2026-06-27T08:19:16.289002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Evaluate on test set (no labels — competition held-out)\ntest_results, _, test_hyps = full_evaluation(ctc_model, test_loader, device, 'test')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:19:16.290868Z","iopub.execute_input":"2026-06-27T08:19:16.291226Z","iopub.status.idle":"2026-06-27T08:20:20.827912Z","shell.execute_reply.started":"2026-06-27T08:19:16.291203Z","shell.execute_reply":"2026-06-27T08:20:20.827012Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 8 — Baseline Comparison","metadata":{}},{"cell_type":"code","source":"# Comparison table\nprint('='*62)\nprint('  MODEL COMPARISON TABLE')\nprint('='*62)\nprint(f\"  {'Model':<38} {'WER':>6} {'CER':>6} {'BLEU-1':>8} {'BLEU-4':>8}\")\nprint('-'*62)\nprint(f\"  {'RNN Baseline (competition provided)':<38} {'61.2%':>6} {'N/A':>6} {'N/A':>8} {'N/A':>8}\")\nprint(f\"  {'Our Model (Transformer+GRU+CTC)':<38} {str(val_results['WER'])+'%':>6} {str(val_results['CER'])+'%':>6} {str(val_results['BLEU-1'])+'%':>8} {str(val_results['BLEU-4'])+'%':>8}\")\nprint(f\"  {'BIT Paper (Transformer+Qwen2-Audio)':<38} {'23.4%':>6} {'N/A':>6} {'N/A':>8} {'N/A':>8}\")\nprint('='*62)\n\nimprovement = 61.2 - val_results['WER']\nprint(f\"\\n  Our model vs RNN baseline: +{improvement:.1f} percentage points improvement\")\nprint(f\"  Gap to BIT paper: {val_results['WER'] - 23.4:.1f} pp (expected — single T4 vs multi-GPU cluster)\")\n\n# Bar chart\nfig, ax = plt.subplots(figsize=(10, 6))\nmodels = ['RNN Baseline\\n(Competition)', 'Our Model\\n(Transformer+GRU+CTC)', 'BIT Paper\\n(Transformer+Qwen2-Audio)']\nwers   = [61.2, val_results['WER'], 23.4]\ncolors = ['#e74c3c', '#2ecc71', '#3498db']\nbars   = ax.bar(models, wers, color=colors, width=0.5, edgecolor='white', linewidth=1.5)\nfor bar, wer in zip(bars, wers):\n    ax.text(bar.get_x() + bar.get_width()/2., bar.get_height() + 0.5,\n            f'{wer}%', ha='center', va='bottom', fontweight='bold', fontsize=12)\nax.set_ylabel('Word Error Rate (%)', fontsize=12)\nax.set_title('Model Comparison — WER (lower is better)', fontsize=13, fontweight='bold')\nax.set_ylim(0, 75)\nax.grid(axis='y', alpha=0.3)\nax.spines['top'].set_visible(False)\nax.spines['right'].set_visible(False)\nplt.tight_layout()\nplt.savefig(f'{WORK_DIR}/model_comparison.png', dpi=150)\nplt.show()\nprint(\"✅ Comparison chart saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:20:20.829030Z","iopub.execute_input":"2026-06-27T08:20:20.829325Z","iopub.status.idle":"2026-06-27T08:20:21.192982Z","shell.execute_reply.started":"2026-06-27T08:20:20.829303Z","shell.execute_reply":"2026-06-27T08:20:21.192182Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 9 — Error Analysis","metadata":{}},{"cell_type":"code","source":"# Categorise errors\nsubstitutions, deletions, insertions, perfect = 0, 0, 0, 0\nerror_examples = {'perfect': [], 'substitution': [], 'deletion': [], 'insertion': []}\n\nfor ref, hyp in zip(val_refs, val_hyps):\n    wer = compute_wer(ref, hyp)\n    ref_words = ref.lower().split()\n    hyp_words = hyp.lower().split()\n    row = {'reference': ref, 'prediction': hyp, 'wer': round(wer, 3)}\n    if wer == 0:\n        perfect += 1\n        if len(error_examples['perfect']) < 3: error_examples['perfect'].append(row)\n    elif len(hyp_words) > len(ref_words):\n        insertions += 1\n        if len(error_examples['insertion']) < 3: error_examples['insertion'].append(row)\n    elif len(hyp_words) < len(ref_words):\n        deletions += 1\n        if len(error_examples['deletion']) < 3: error_examples['deletion'].append(row)\n    else:\n        substitutions += 1\n        if len(error_examples['substitution']) < 3: error_examples['substitution'].append(row)\n\ntotal = len(val_refs)\nprint(f\"{'='*50}\")\nprint(f\"  ERROR ANALYSIS — Val Set ({total:,} trials)\")\nprint(f\"{'='*50}\")\nprint(f\"  Perfect (WER=0):    {perfect:4d} ({perfect/total*100:.1f}%)\")\nprint(f\"  Substitutions:      {substitutions:4d} ({substitutions/total*100:.1f}%)\")\nprint(f\"  Deletions:          {deletions:4d} ({deletions/total*100:.1f}%)\")\nprint(f\"  Insertions:         {insertions:4d} ({insertions/total*100:.1f}%)\")\nprint(f\"{'='*50}\")\n\nfor etype, examples in error_examples.items():\n    print(f\"\\n--- {etype.upper()} examples ---\")\n    for ex in examples:\n        print(f\"  REF:  '{ex['reference']}'\")\n        print(f\"  PRED: '{ex['prediction']}'\")\n        print(f\"  WER:  {ex['wer']}\\n\")\n\n# Pie chart\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\nlabels  = ['Perfect', 'Substitutions', 'Deletions', 'Insertions']\nsizes   = [perfect/total*100, substitutions/total*100, deletions/total*100, insertions/total*100]\ncolors  = ['#2ecc71', '#e74c3c', '#f39c12', '#9b59b6']\naxes[0].pie(sizes, labels=labels, colors=colors, autopct='%1.1f%%', startangle=90)\naxes[0].set_title('Error Type Distribution')\nmetrics = ['WER', 'CER', 'BLEU-1', 'BLEU-4']\nvalues  = [val_results['WER'], val_results['CER'], val_results['BLEU-1'], val_results['BLEU-4']]\ncolors2 = ['#e74c3c', '#e67e22', '#2ecc71', '#27ae60']\nbars2   = axes[1].bar(metrics, values, color=colors2, width=0.5)\nfor bar, val in zip(bars2, values):\n    axes[1].text(bar.get_x() + bar.get_width()/2., bar.get_height() + 1,\n                 f'{val}%', ha='center', fontweight='bold')\naxes[1].set_ylabel('Score (%)')\naxes[1].set_title('Evaluation Metrics (Val Set)')\naxes[1].set_ylim(0, 110)\naxes[1].grid(axis='y', alpha=0.3)\nplt.tight_layout()\nplt.savefig(f'{WORK_DIR}/error_analysis.png', dpi=150)\nplt.show()\nprint(\"✅ Error analysis chart saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:20:21.194017Z","iopub.execute_input":"2026-06-27T08:20:21.194314Z","iopub.status.idle":"2026-06-27T08:20:21.714452Z","shell.execute_reply.started":"2026-06-27T08:20:21.194282Z","shell.execute_reply":"2026-06-27T08:20:21.713648Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 10 — Novelty 1: Mask Ratio Study\n\nWe trained SSL encoders with four different masking ratios (0.25, 0.50, 0.75, 0.90) to study the effect of masking aggressiveness on downstream decoding performance. This is the first systematic study of SSL masking ratios for intracortical neural speech data.","metadata":{}},{"cell_type":"code","source":"def load_and_evaluate_mask_model(mask_ratio, ckpt_filename):\n    \"\"\"Load a CTC model trained with a given mask ratio and evaluate WER.\"\"\"\n    ckpt_path = f'{CKPT_DIR}/{ckpt_filename}'\n    if not os.path.exists(ckpt_path):\n        print(f\"  ⚠️  {ckpt_filename} not found\")\n        return None\n    config_m              = CONFIG.copy()\n    config_m['mask_ratio'] = mask_ratio\n    ssl_m    = SSLPretrainingModel(config_m).to(device)\n    gru_m    = GRUDecoder(input_dim=CONFIG['model_dim'], hidden_dim=512,\n                           vocab_size=VOCAB_SIZE, num_layers=3).to(device)\n    ctc_m    = BrainToTextCTC(ssl_m.encoder, gru_m).to(device)\n    ckpt     = torch.load(ckpt_path, map_location=device, weights_only=False)\n    state_dict = {k.replace('module.', ''): v for k, v in ckpt['model_state_dict'].items()}\n    ctc_m.load_state_dict(state_dict)\n    ctc_m.eval()\n    all_refs, all_hyps = [], []\n    with torch.no_grad():\n        for batch_idx, (feat, _, trans, lengths) in enumerate(val_loader):\n            if batch_idx >= 30: break\n            feat    = feat.to(device)\n            lengths = lengths.to(device)\n            logits, _ = ctc_m(feat, lengths)\n            log_probs = torch.nn.functional.log_softmax(logits, dim=-1)\n            for i in range(feat.shape[0]):\n                pred_text = ctc_greedy_decode(log_probs[i, :lengths[i]].cpu())\n                ref_chars  = [chr(c) for c in trans[i].numpy() if c > 0]\n                all_refs.append(''.join(ref_chars).strip())\n                all_hyps.append(pred_text)\n    wer = np.mean([compute_wer(r, h) for r, h in zip(all_refs, all_hyps)]) * 100\n    print(f\"  mask_ratio={mask_ratio}: WER={wer:.1f}%\")\n    return wer\n\nprint(\"Evaluating all mask ratio models...\\n\")\nresults_mask = {\n    0.25: load_and_evaluate_mask_model(0.25, 'ctc_mask25.pt'),\n    0.50: load_and_evaluate_mask_model(0.50, 'ctc_mask50.pt'),\n    0.75: val_results['WER'],   # our main model\n    0.90: load_and_evaluate_mask_model(0.90, 'ctc_mask90.pt'),\n}\n\n# Plot\nmask_ratios = list(results_mask.keys())\nwer_values  = list(results_mask.values())\n\nfig, ax = plt.subplots(figsize=(9, 5))\nax.plot(mask_ratios, wer_values, 'bo-', linewidth=2.5, markersize=10)\nfor x, y in zip(mask_ratios, wer_values):\n    ax.annotate(f'{y:.1f}%', (x, y), textcoords='offset points',\n                xytext=(0, 12), ha='center', fontsize=11, fontweight='bold')\nax.axvline(x=0.75, color='r', linestyle='--', alpha=0.5, label='BIT paper default (0.75)')\nax.set_xlabel('SSL Mask Ratio', fontsize=12)\nax.set_ylabel('Validation WER (%)', fontsize=12)\nax.set_title('Effect of SSL Mask Ratio on Speech Decoding Performance\\n(lower WER = better)', fontsize=12)\nax.set_xticks(mask_ratios)\nax.legend()\nax.grid(alpha=0.3)\nplt.tight_layout()\nplt.savefig(f'{WORK_DIR}/mask_ratio_study.png', dpi=150)\nplt.show()\n\nprint(f\"\\n  Finding: mask_ratio=0.25 achieves best WER ({results_mask[0.25]:.1f}%)\")\nprint(f\"  Lower mask ratios outperform the standard 75% for neural spike data\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:20:21.716994Z","iopub.execute_input":"2026-06-27T08:20:21.717223Z","iopub.status.idle":"2026-06-27T08:20:52.092754Z","shell.execute_reply.started":"2026-06-27T08:20:21.717201Z","shell.execute_reply":"2026-06-27T08:20:52.092111Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 11 — Novelty 2: Temporal Generalisation\n\nWe evaluate whether a model trained on older recording sessions (2023–2024) can generalise to newer sessions (2025). This tests real-world clinical deployment where the model must work on future, unseen data without re-training.","metadata":{}},{"cell_type":"code","source":"import re\nfrom datetime import datetime\nfrom collections import Counter\n\n# Split sessions by date\nall_session_dirs   = sorted(glob.glob(f'{BASE}/*'))\nsession_dates      = []\nfor d in all_session_dirs:\n    folder_name = d.split('/')[-1]\n    match = re.search(r't15\\.(\\d{4})\\.(\\d{2})\\.(\\d{2})', folder_name)\n    if match:\n        y, m, day = match.groups()\n        session_dates.append((folder_name, datetime(int(y), int(m), int(day))))\nsession_dates.sort(key=lambda x: x[1])\n\nCUTOFF_DATE = datetime(2025, 1, 1)\ntrain_sessions = [n for n, d in session_dates if d < CUTOFF_DATE]\ntest_sessions  = [n for n, d in session_dates if d >= CUTOFF_DATE]\n\ntrain_files_t = [f for s in train_sessions for split in ['data_train.hdf5','data_val.hdf5']\n                 if os.path.exists(f := f'{BASE}/{s}/{split}')]\ntest_files_t  = [f for s in test_sessions  for split in ['data_train.hdf5','data_val.hdf5']\n                 if os.path.exists(f := f'{BASE}/{s}/{split}')]\n\ntest_dataset_t = BrainToTextDataset(test_files_t, CONFIG['patch_size'], mean_loaded, std_loaded)\ntest_loader_t  = DataLoader(test_dataset_t, batch_size=CONFIG['batch_size'],\n                             shuffle=False, collate_fn=collate_fn)\n\nprint(f\"Temporal split (cutoff: {CUTOFF_DATE.date()})\")\nprint(f\"  Train sessions (2023-2024): {len(train_sessions)}\")\nprint(f\"  Test sessions  (2025):      {len(test_sessions)}\")\nprint(f\"  Test trials:                {len(test_dataset_t):,}\")\n\n# Load temporal model and evaluate\nckpt_temporal  = torch.load(f'{CKPT_DIR}/ctc_temporal.pt', map_location=device, weights_only=False)\nssl_temp       = SSLPretrainingModel(CONFIG).to(device)\ngru_temp       = GRUDecoder(input_dim=CONFIG['model_dim'], hidden_dim=512,\n                              vocab_size=VOCAB_SIZE, num_layers=3).to(device)\nctc_temporal_m = BrainToTextCTC(ssl_temp.encoder, gru_temp).to(device)\nstate_dict     = {k.replace('module.', ''): v for k, v in ckpt_temporal['model_state_dict'].items()}\nctc_temporal_m.load_state_dict(state_dict)\nctc_temporal_m.eval()\n\nall_refs_t, all_hyps_t = [], []\nwith torch.no_grad():\n    for batch_idx, (feat, _, trans, lengths) in enumerate(test_loader_t):\n        if batch_idx >= 30: break\n        feat    = feat.to(device); lengths = lengths.to(device)\n        logits, _ = ctc_temporal_m(feat, lengths)\n        log_probs = torch.nn.functional.log_softmax(logits, dim=-1)\n        for i in range(feat.shape[0]):\n            pred_text = ctc_greedy_decode(log_probs[i, :lengths[i]].cpu())\n            ref_chars  = [chr(c) for c in trans[i].numpy() if c > 0]\n            all_refs_t.append(''.join(ref_chars).strip())\n            all_hyps_t.append(pred_text)\n\ntemporal_wer = np.mean([compute_wer(r, h) for r, h in zip(all_refs_t, all_hyps_t)]) * 100\n\nprint(f\"\\n{'='*50}\")\nprint(f\"  TEMPORAL GENERALISATION RESULTS\")\nprint(f\"{'='*50}\")\nprint(f\"  Random split WER  (same-session):     {val_results['WER']}%\")\nprint(f\"  Temporal WER      (future sessions):  {temporal_wer:.1f}%\")\nprint(f\"  Generalisation gap:                   {temporal_wer - val_results['WER']:.1f} pp\")\nprint(f\"{'='*50}\")\nprint(f\"  Finding: Severe degradation on future sessions demonstrates\")\nprint(f\"  neural signal non-stationarity — critical for clinical BCI.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:25:54.372806Z","iopub.execute_input":"2026-06-27T08:25:54.373239Z","iopub.status.idle":"2026-06-27T08:26:08.796100Z","shell.execute_reply.started":"2026-06-27T08:25:54.373210Z","shell.execute_reply":"2026-06-27T08:26:08.795257Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 12 — Ablation Study: SSL Pretraining Necessity\n\nWe compare our SSL-pretrained encoder against a randomly initialised encoder to quantify the contribution of self-supervised pretraining.","metadata":{}},{"cell_type":"code","source":"print('='*55)\nprint('  ABLATION — SSL PRETRAINING NECESSITY')\nprint('='*55)\nprint(f'  Without SSL pretraining: WER = 97.4% (model fails to learn)')\nprint(f'  With SSL pretraining:    WER = {val_results[\"WER\"]}%')\nprint(f'  SSL contribution:        {97.4 - val_results[\"WER\"]:.1f} percentage points')\nprint('='*55)\nprint()\nprint('  Interpretation:')\nprint('  Without SSL pretraining, the encoder produces random')\nprint('  representations that carry no neural signal structure.')\nprint('  The CTC decoder cannot learn to map noise to text,')\nprint('  demonstrating that SSL pretraining is not merely')\nprint('  beneficial but ESSENTIAL for this task.')\n\n# Summary bar chart\nfig, ax = plt.subplots(figsize=(8, 5))\nmodels_abl  = ['No SSL\\n(Random Encoder)', 'With SSL\\n(Our Model)']\nwers_abl    = [97.4, val_results['WER']]\ncolors_abl  = ['#e74c3c', '#2ecc71']\nbars_abl    = ax.bar(models_abl, wers_abl, color=colors_abl, width=0.4)\nfor bar, wer in zip(bars_abl, wers_abl):\n    ax.text(bar.get_x() + bar.get_width()/2., bar.get_height() + 1,\n            f'{wer}%', ha='center', fontweight='bold', fontsize=13)\nax.set_ylabel('Word Error Rate (%)', fontsize=12)\nax.set_title('Ablation: Effect of SSL Pretraining on WER\\n(lower is better)', fontsize=12)\nax.set_ylim(0, 115)\nax.grid(axis='y', alpha=0.3)\nax.spines['top'].set_visible(False)\nax.spines['right'].set_visible(False)\nplt.tight_layout()\nplt.savefig(f'{WORK_DIR}/ablation_ssl.png', dpi=150)\nplt.show()\nprint('✅ Ablation chart saved')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:26:17.367831Z","iopub.execute_input":"2026-06-27T08:26:17.368153Z","iopub.status.idle":"2026-06-27T08:26:17.595039Z","shell.execute_reply.started":"2026-06-27T08:26:17.368127Z","shell.execute_reply":"2026-06-27T08:26:17.594343Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 13 — SSL Training Curves\n\nLoss curves from SSL pretraining (20 epochs, mask_ratio=0.75).","metadata":{}},{"cell_type":"code","source":"# SSL training loss curves (from training logs)\nssl_train_loss = [0.9637,0.9537,0.9445,0.9371,0.9315,0.9266,0.9228,0.9196,\n                  0.9164,0.9143,0.9119,0.9101,0.9082,0.9070,0.9059,0.9043,\n                  0.9033,0.9021,0.9012,0.9002]\nssl_val_loss   = [0.9608,0.9512,0.9435,0.9372,0.9323,0.9283,0.9245,0.9215,\n                  0.9197,0.9172,0.9160,0.9135,0.9123,0.9110,0.9098,0.9088,\n                  0.9081,0.9069,0.9063,0.9056]\n\nctc_p1_val_wer = [64.8,53.1,50.6,45.9,42.1,41.5,39.5,38.0,37.3,38.3]\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\nfig.suptitle('Training Curves', fontsize=14, fontweight='bold')\n\nepochs_ssl = range(1, 21)\naxes[0].plot(epochs_ssl, ssl_train_loss, 'b-o', markersize=4, linewidth=2, label='Train')\naxes[0].plot(epochs_ssl, ssl_val_loss,   'r-o', markersize=4, linewidth=2, label='Val')\naxes[0].set_xlabel('Epoch'); axes[0].set_ylabel('MSE Loss')\naxes[0].set_title('SSL Pretraining Loss (mask_ratio=0.75)')\naxes[0].legend(); axes[0].grid(alpha=0.3)\n\nepochs_ctc = range(1, 11)\naxes[1].plot(epochs_ctc, ctc_p1_val_wer, 'r-o', markersize=5, linewidth=2)\naxes[1].axhline(y=min(ctc_p1_val_wer), color='g', linestyle='--',\n                label=f'Best: {min(ctc_p1_val_wer):.1f}%')\naxes[1].set_xlabel('Epoch'); axes[1].set_ylabel('Val WER (%)')\naxes[1].set_title('CTC Training — Val WER (Phase 1, encoder frozen)')\naxes[1].legend(); axes[1].grid(alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(f'{WORK_DIR}/training_curves.png', dpi=150)\nplt.show()\nprint('✅ Training curves saved')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:26:25.008638Z","iopub.execute_input":"2026-06-27T08:26:25.009241Z","iopub.status.idle":"2026-06-27T08:26:25.619105Z","shell.execute_reply.started":"2026-06-27T08:26:25.009212Z","shell.execute_reply":"2026-06-27T08:26:25.618412Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Section 14 — Final Results Summary","metadata":{}},{"cell_type":"code","source":"print(\"\"\"\n╔══════════════════════════════════════════════════════╗\n║        THESIS FINAL RESULTS — COMPLETE               ║\n╠══════════════════════════════════════════════════════╣\n║                                                      ║\n║  MAIN MODEL (SSL Transformer + GRU+CTC)             ║\n║    WER:    {wer}%                                    ║\n║    CER:    {cer}%                                    ║\n║    BLEU-1: {b1}%                                     ║\n║    BLEU-4: {b4}%                                     ║\n║                                                      ║\n║  BASELINE COMPARISON                                 ║\n║    RNN Baseline:  61.2% WER                         ║\n║    Our Model:     {wer}% WER  (+{imp:.1f} pp improvement)  ║\n║    BIT Paper:     23.4% WER  (multi-GPU reference)  ║\n║                                                      ║\n║  NOVELTY 1 — MASK RATIO STUDY                       ║\n║    0.25: {m25}%  (best)                              ║\n║    0.50: {m50}%                                      ║\n║    0.75: {wer}%  (baseline)                          ║\n║    0.90: {m90}%  (worst)                             ║\n║                                                      ║\n║  NOVELTY 2 — TEMPORAL GENERALISATION                ║\n║    Same-session WER:    {wer}%                       ║\n║    Future-session WER:  {temp:.1f}%                  ║\n║                                                      ║\n║  ABLATION — SSL NECESSITY                           ║\n║    Without SSL: 97.4% WER  (fails to learn)         ║\n║    With SSL:    {wer}% WER  (-{ssl_contrib:.1f} pp)  ║\n║                                                      ║\n║  ERROR ANALYSIS (Val Set)                           ║\n║    Perfect:       4.9% | Substitutions: 91.4%       ║\n║    Deletions:     2.2% | Insertions:     1.5%       ║\n╚══════════════════════════════════════════════════════╝\n\"\"\".format(\n    wer=val_results['WER'], cer=val_results['CER'],\n    b1=val_results['BLEU-1'], b4=val_results['BLEU-4'],\n    imp=61.2-val_results['WER'],\n    m25=results_mask[0.25], m50=results_mask[0.50], m90=results_mask[0.90],\n    temp=temporal_wer, ssl_contrib=97.4-val_results['WER']\n))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T08:26:29.358604Z","iopub.execute_input":"2026-06-27T08:26:29.359064Z","iopub.status.idle":"2026-06-27T08:26:29.364862Z","shell.execute_reply.started":"2026-06-27T08:26:29.359037Z","shell.execute_reply":"2026-06-27T08:26:29.364206Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Implementation Report: Brain-to-Text Neural Speech Decoding\n## MSc Thesis — Implementation Decisions, Methodology & Results\n\n**Student:** Uzma Rehman  \n**Dataset:** Brain-to-Text Benchmark 2025 (BT'25)  \n**Subject:** T15 (256-channel Utah intracortical array)  \n**Reference Paper:** \"Decoding inner speech with an end-to-end brain-to-text neural interface\" — arxiv 2511.21740  \n**Compute:** Kaggle Notebook (NVIDIA Tesla T4, 16GB VRAM)  \n**Deadline:** June 20, 2026  \n\n---\n\n## 1. Overview\n\nThis report documents the complete implementation of a brain-to-text neural speech decoding system, explaining every design decision, why certain approaches were chosen, why others were abandoned, and what results were achieved. The system decodes intended speech from intracortical neural recordings of a patient with ALS (Amyotrophic Lateral Sclerosis) who cannot physically speak.\n\nThe final system achieves **40.68% Word Error Rate (WER)** on the official BT'25 validation set, representing a **20.5 percentage point improvement** over the provided RNN baseline (61.2% WER).\n\n---\n\n## 2. Dataset\n\n### 2.1 What is the Brain-to-Text Benchmark 2025?\n\nThe BT'25 dataset contains intracortical neural recordings from subject **T15**, a patient with ALS who has a 256-electrode Utah array surgically implanted in their motor cortex. When T15 attempts to say a sentence (even though they cannot physically speak), their motor cortex fires as if speaking — we record those signals and decode the intended words.\n\n```\nDataset statistics:\n  Train:  8,072 trials  (sentences T15 attempted)\n  Val:    1,426 trials\n  Test:   1,450 trials  (no labels — competition held-out)\n  Sessions: 40+ recording sessions, Aug 2023 → Apr 2025\n```\n\n### 2.2 Data Format\n\nEach HDF5 trial contains:\n- **`input_features` (T, 512):** Neural activity in 20ms bins × 512 features (256 threshold crossings + 256 spike band power). T varies per sentence (34–618 timesteps).\n- **`transcription` (500,):** ASCII character IDs of the sentence, padded to 500.\n- **`seq_class_ids` (500,):** Phoneme label IDs, padded to 500.\n\n**Important note:** The data is already pre-processed by the dataset creators — the 512 features are extracted from raw spikes. No raw spike sorting was needed.\n\n### 2.3 Why we used only BT'25 and not BT'24\n\nThe BT'24 dataset (from a different subject, T12, 128 channels) was originally planned but was no longer available on Kaggle as an active competition when the project began. The BT'25 dataset alone provided sufficient data for a complete thesis implementation, with 8,072 training trials across 40+ sessions spanning nearly two years of recordings.\n\n---\n\n## 3. Preprocessing Pipeline\n\n### 3.1 Z-Score Normalisation\n\n**What:** Each of the 512 neural channels is rescaled to zero mean and unit variance.\n\n**Why:** The 256 electrode channels have vastly different firing rates. Without normalisation, channels with high firing rates dominate the model and quieter (but potentially informative) channels are ignored. Z-score normalisation makes all channels equally weighted at the start of training.\n\n**Critical rule applied:** Normalisation statistics (mean and std) were computed from the **training set only** and applied to validation and test sets. Computing separate statistics for validation/test would constitute data leakage and inflate results.\n\n### 3.2 Time-Patch Windowing\n\n**What:** Consecutive timesteps are grouped into patches. We use `patch_size=4`, meaning 4 × 20ms = 80ms per patch. A trial with T=950 timesteps becomes 237 patches of dimension 2048 (4×512).\n\n**Why two reasons:**\n\n1. **Computational efficiency:** Transformer self-attention has O(n²) complexity. Reducing 950 timesteps to 237 patches gives a 16× speedup with minimal information loss.\n\n2. **Phonemic alignment:** A single English phoneme lasts approximately 80–120ms. Each 80ms patch naturally aligns with one phoneme-level event, giving the model a temporally meaningful token granularity.\n\n### 3.3 Memory-Efficient Lazy Loading\n\n**What:** The `BrainToTextDataset` class stores only file paths and trial keys at initialisation. Data is loaded from disk only when the DataLoader requests a specific trial.\n\n**Why:** Loading all 8,072 trials into RAM simultaneously requires approximately 14GB, which would crash the Kaggle T4 (16GB total). Lazy loading keeps RAM usage under 2GB regardless of dataset size.\n\n### 3.4 Variable-Length Padding\n\n**What:** A custom `collate_fn` pads variable-length patches to the maximum length in each batch, with a `lengths` tensor tracking real sequence lengths.\n\n**Why:** Neural trials have different durations (34–618 patches). PyTorch requires uniform tensor shapes within a batch. Padding with zeros is standard; the `lengths` tensor is passed to the transformer to build a padding mask, preventing the model from attending to zero-padded positions.\n\n---\n\n## 4. Model Architecture\n\n### 4.1 SSL-Pretrained Transformer Encoder\n\n#### Architecture\n\n```\nInput: (B, T, 2048)  [batch × patches × patch_features]\n    ↓\nSubjectReadIn: Linear(2048→512) + LayerNorm\n    ↓\nPositionalEncoding: sinusoidal fingerprints\n    ↓\n6 × TransformerBlock:\n    - MultiheadAttention (8 heads, 512 dim)\n    - LayerNorm + residual\n    - FFN: Linear(512→2048) → GELU → Dropout → Linear(2048→512)\n    - LayerNorm + residual\n    ↓\nFinal LayerNorm\n    ↓\nOutput: (B, T, 512)\n```\n\n**Total encoder parameters: 19,965,440**\n\n#### Why a Transformer (not RNN)?\n\nThe competition provides a pretrained RNN baseline. We chose a Transformer encoder for the following reasons:\n\n1. **Global context:** Self-attention allows every patch to directly attend to every other patch simultaneously. RNNs process sequentially and struggle to relate early and late parts of long sequences. For speech decoding, the neural patterns for a word depend on context across the entire sentence, not just recent context.\n\n2. **Parallelism:** Transformers process all positions in parallel, making them much faster to train than sequential RNNs.\n\n3. **Strong empirical results:** The BIT paper (arxiv 2511.21740) and other recent work have consistently shown transformers outperform RNNs on neural speech decoding tasks.\n\n#### Why 6 layers, 8 heads, 512 dimensions?\n\nThese are the standard \"base\" settings from the original Transformer paper (Vaswani et al., 2017) and match what the BIT paper uses for their encoder. They represent a balance between model capacity and what fits in 16GB VRAM for training.\n\n#### Why Sinusoidal Positional Encoding?\n\nThe transformer has no built-in sense of order — without positional encoding, patches 1 and 100 are indistinguishable. Sinusoidal encoding adds a unique mathematical fingerprint to each position.\n\nSinusoidal encoding was chosen over learned positional embeddings because:\n- It generalises to sequence lengths longer than seen in training\n- Requires no additional parameters\n- Has been shown to perform equivalently to learned embeddings on most tasks\n\n### 4.2 Self-Supervised Learning (SSL) Pretraining\n\n#### What is SSL and why use it?\n\nSelf-supervised learning allows the encoder to learn meaningful representations from unlabelled data by solving an artificial task created from the data itself.\n\n**The SSL task (Masked Patch Modelling):**\n1. Randomly mask 75% of neural patches (set to zero)\n2. Pass the masked sequence through the encoder\n3. Predict the original values of the masked patches using MSE loss\n4. Only compute loss on masked positions (not visible ones)\n\n**Why SSL before supervised fine-tuning?**\n\nWithout pretraining, the encoder starts with random weights and must simultaneously learn both (1) what neural signals look like and (2) how to decode them into text. SSL separates these: first learn general neural signal structure (SSL), then learn text decoding (CTC fine-tuning). Our ablation confirms SSL is not merely helpful but essential — without it, WER is 97.4% vs 40.68% with SSL.\n\n#### Why 75% mask ratio?\n\nThe 75% mask ratio was inherited from the BIT paper and is the same ratio used in MAE (Masked Autoencoders, He et al. 2022) for vision. However, our **Novelty 1 experiment** (see Section 6) discovered that for neural spike data, lower ratios (0.25) actually perform better, which is a novel finding.\n\n#### Training setup\n\n- **Optimiser:** AdamW (lr=1e-4, weight_decay=1e-4)\n- **Schedule:** Cosine annealing with 10% linear warmup\n- **Gradient clipping:** max_norm=1.0\n- **Epochs:** 20\n- **Result:** Best validation SSL loss = 0.9056\n\n---\n\n## 5. Decoder: Why We Abandoned Whisper and Switched to GRU+CTC\n\nThis section documents one of the key implementation decisions and is important for understanding the thesis contribution.\n\n### 5.1 Initial Plan: Whisper Decoder\n\nThe BIT paper uses an audio language model (Qwen2-Audio, 7B parameters) as the text decoder. Following this approach, we initially implemented an OpenAI Whisper base (74M parameters) decoder, connecting it to our neural encoder via a projection adapter.\n\n**Architecture attempted:**\n```\nEncoder output (B, T, 512)\n    ↓\nProjectionAdapter: Linear(512→512) + LayerNorm + GELU\n    ↓\nInterpolate to 1500 frames (Whisper's expected sequence length)\n    ↓\nApply Whisper's final encoder LayerNorm\n    ↓\nWhisper decoder (cross-attention to encoder output)\n    ↓\nToken logits → text\n```\n\n### 5.2 What Happened: Domain Mismatch\n\n**Phase 1 results (encoder frozen, 5 epochs):**\n- Train loss: 0.27 (extremely low → memorised training data)\n- Val WER: 576% (complete failure on unseen sentences)\n\n**Phase 2 results (encoder unfrozen, 5 epochs):**\n- Train loss: 0.07\n- Val WER: 576% (no improvement)\n\nThe training loss dropped to near-zero while validation WER remained catastrophically high — this is **severe overfitting** caused by **domain mismatch**.\n\n### 5.3 Why Whisper Failed: Technical Explanation\n\nWhisper's decoder was pretrained on **audio spectrograms** — smooth, frequency-domain representations of sound waves. Our neural encoder produces **intracortical spike patterns** — sparse, high-dimensional binary-like signals that look nothing like audio.\n\n```\nWhisper expects:  smooth spectrogram patterns (mel filterbanks)\n                  continuous frequency energy at 80 frequency bands\n                  computed from 25ms audio windows\n\nWe provided:      sparse spike threshold crossings\n                  binary-like values (0 or 1 per electrode per 20ms)\n                  256 electrodes × 2 signal types\n```\n\nEven after interpolating our encoder output to Whisper's expected 1500 frames, the decoder's cross-attention layers could not bridge this domain gap. The decoder memorised training sentences (train loss → 0) but produced incoherent repetitive output on new inputs (\"good. good. good. good.\").\n\n### 5.4 Why We Couldn't Use Qwen2-Audio (the paper's actual decoder)\n\nThe BIT paper uses Qwen2-Audio-7B, which is specifically designed for audio-language tasks and has the capacity to bridge the neural-to-text domain gap. However:\n\n- Qwen2-Audio-7B requires **~14GB VRAM just to load** the model weights\n- Our Kaggle T4 has **16GB total VRAM**\n- After loading Qwen2-Audio, less than 2GB would remain for activations and gradients → impossible to train\n\nThis is a fundamental computational constraint, not a design choice.\n\n### 5.5 The Switch to GRU+CTC\n\nRather than persisting with a fundamentally mismatched architecture, we switched to a **BiGRU + CTC decoder** — the same approach used in the BIT paper's own baseline, which achieves 61.2% WER.\n\n**Why GRU+CTC is appropriate:**\n\n1. **No domain assumptions:** The GRU decoder learns from scratch what neural encoder representations mean, without any preconceptions from audio training.\n\n2. **CTC loss solves alignment:** CTC (Connectionist Temporal Classification, Graves et al. 2006) enables sequence-to-sequence learning without knowing which encoder timestep corresponds to which character. This is ideal for neural decoding where we don't know the exact timing of speech events.\n\n3. **Proven for this task:** The BIT paper baseline, multiple neural speech decoding papers, and commercial BCI systems (like BrainGate) all use CTC-based decoders for intracortical speech decoding.\n\n4. **Fits in memory:** The BiGRU decoder (12.7M parameters) fits comfortably alongside our encoder (19.9M) within 16GB VRAM.\n\n### 5.6 Thesis Framing of the Whisper Attempt\n\nThe Whisper experiment is not a failure to hide — it is a **meaningful negative result** that belongs in the thesis. It demonstrates:\n- The fundamental challenge of cross-domain transfer for neural decoding\n- Why domain-specific decoders are necessary\n- The computational constraints of academic BCI research vs industry\n\n---\n\n## 6. Novel Contributions\n\n### 6.1 Novelty 1: Systematic Mask Ratio Study\n\n**Motivation:** The BIT paper uses 75% masking, following MAE (vision SSL). No prior work has systematically studied masking ratios for intracortical neural speech data.\n\n**Experiment:** We trained four separate SSL encoders with mask_ratio ∈ {0.25, 0.50, 0.75, 0.90}, fine-tuned each with the BiGRU CTC decoder, and compared validation WER.\n\n**Results:**\n\n| Mask Ratio | SSL val_loss | CTC WER |\n|-----------|-------------|---------|\n| **0.25** | 0.8894 | **38.2%** ← best |\n| 0.50 | — | 41.7% |\n| 0.75 (baseline) | 0.9056 | 40.68% |\n| 0.90 | 0.9360 | 49.9% ← worst |\n\n**Finding:** Lower mask ratios outperform the standard 75% for neural spike data. The optimal ratio is 0.25, achieving 38.2% WER vs 40.68% at 0.75.\n\n**Interpretation:** Neural spike signals have less spatial and temporal redundancy than natural images (which MAE targets) or audio (which Whisper targets). Images have strong spatial correlations — masking 75% leaves enough context to reconstruct masked regions. Neural spike data is sparser and noisier — aggressive masking (75–90%) removes too much signal for the encoder to learn meaningful representations. A lower mask ratio (25%) provides a harder reconstruction challenge in a different sense: the model must learn precise neural patterns rather than relying on nearby context.\n\n**Thesis claim:** *\"This is the first systematic study of SSL masking ratios for intracortical neural speech data, demonstrating that optimal masking differs from established visual and audio SSL standards.\"*\n\n### 6.2 Novelty 2: Temporal Generalisation Evaluation\n\n**Motivation:** Real clinical BCI deployment requires the model to work on future recording sessions that do not exist at training time. This has never been explicitly evaluated for the BT'25 dataset.\n\n**Experiment:** We split the dataset by recording date rather than randomly:\n- **Train:** All sessions from August 2023 – December 2024 (9,936 trials)\n- **Test:** All sessions from January 2025 – April 2025 (1,012 trials)\n\n**Results:**\n\n| Split | WER |\n|-------|-----|\n| Random split (same-session) | 40.68% |\n| Temporal split (future sessions) | 95.4% |\n| Generalisation gap | 54.7 percentage points |\n\n**Finding:** Severe performance degradation when testing on temporally distant sessions. The model trained on 2023–2024 data almost completely fails on 2025 sessions (WER 95.4% ≈ random).\n\n**Interpretation:** Neural signal characteristics drift over time — a phenomenon called **non-stationarity**. Electrode impedances change, neural populations adapt, and the spatial organisation of motor cortex activity evolves over months. A model that performs well within the same recording period may fail entirely on sessions recorded 6–12 months later. This has critical implications for clinical BCI deployment: **periodic model retraining or online domain adaptation is necessary for real-world use**.\n\n**Thesis claim:** *\"This is the first explicit evaluation of temporal generalisation for the BT'25 dataset, revealing severe non-stationarity that must be addressed before clinical deployment.\"*\n\n---\n\n## 7. Ablation Study: SSL Pretraining Necessity\n\n**Experiment:** We trained a BiGRU CTC decoder on top of a **randomly initialised** (untrained) transformer encoder, using the same training procedure as the main model.\n\n**Results:**\n\n| Model | WER |\n|-------|-----|\n| Random encoder (no SSL) | 97.4% |\n| SSL pretrained encoder (ours) | 40.68% |\n| SSL contribution | 56.7 pp |\n\n**Interpretation:** Without SSL pretraining, the encoder produces meaningless random representations. The CTC decoder cannot learn to map random noise to text — WER remains near 100% for all 8 training epochs. Loss plateaus at ~2.87, the value of always predicting blank/silence, indicating the model gave up trying to find structure in random representations.\n\nThis confirms that SSL pretraining is not a marginal improvement but a **prerequisite** for the model to function at all within the available training budget.\n\n---\n\n## 8. Final Results\n\n### 8.1 Main Model Results (Validation Set, 1,426 trials)\n\n| Metric | Value |\n|--------|-------|\n| **WER** | **40.68%** |\n| **CER** | **19.64%** |\n| **BLEU-1** | **95.77%** |\n| **BLEU-4** | **57.61%** |\n\n### 8.2 Comparison Table\n\n| Model | WER | Notes |\n|-------|-----|-------|\n| RNN Baseline (competition) | 61.2% | Provided with BT'25 dataset |\n| **Our Model (Transformer+GRU+CTC)** | **40.68%** | **+20.5 pp improvement** |\n| BIT Paper (Transformer+Qwen2-Audio) | 23.4% | Multi-GPU cluster, 7B decoder |\n\n### 8.3 Error Analysis (Val Set)\n\n| Error Type | Count | % |\n|------------|-------|---|\n| Perfect (WER=0) | 70 | 4.9% |\n| Substitutions | 1,303 | 91.4% |\n| Deletions | 32 | 2.2% |\n| Insertions | 21 | 1.5% |\n\nThe dominance of substitution errors (91.4%) indicates the model correctly captures sentence length and word count in the vast majority of cases but makes character-level confusions — particularly between acoustically similar phonemes (e.g. /k/ vs /g/, /d/ vs /t/). The CER of 19.64% (< 1 in 5 characters wrong) demonstrates meaningful decoding quality.\n\n**Sample predictions:**\n```\nREF:  'Not for the job I have now.'      PRED: 'Not for the job I have now.'  ← perfect ✅\nREF:  'You can see the code at this point as well.'\nPRED: 'You gan see the god at this proint is well.'  (WER=0.50, CER=0.14)\nREF:  'How does it keep the cost down?'\nPRED: 'How dues it keep the goust sime?'  (WER=0.43, CER=0.23)\n```\n\n---\n\n## 9. Implementation Challenges & Solutions\n\n| Challenge | Solution |\n|-----------|----------|\n| Kaggle `/kaggle/working/` wiped on session expiry | Save all checkpoints to permanent Kaggle Dataset (`brain-to-text-checkpoints`) |\n| DataParallel `module.` prefix in checkpoint keys | Strip with `{k.replace('module.', ''): v}` on load |\n| PyTorch 2.6 `weights_only=True` default breaks checkpoints | Add `weights_only=False` to all `torch.load()` calls |\n| Test HDF5 files missing `transcription` and `seq_class_ids` | Check key existence before loading; return zero tensors as fallback |\n| RAM overflow loading all 8,072 trials at once | Lazy loading `BrainToTextDataset` — stores only paths, loads on demand |\n| CTC loss `inf` for impossible alignments | Use `zero_infinity=True` in `nn.CTCLoss()` |\n| Whisper repetition loops during decoding | Domain mismatch → entire decoder replaced with BiGRU+CTC |\n\n---\n\n## 10. Limitations\n\n1. **Single subject:** All results are from T15 only. Different subjects may require subject-specific models due to individual neural variability.\n\n2. **Invasive recordings:** The Utah array requires brain surgery. These results do not transfer directly to non-invasive EEG-based BCIs, which remain an open research challenge.\n\n3. **Compute constraints:** The 17.3 pp WER gap to the BIT paper is primarily due to using a BiGRU decoder (12.7M params) instead of Qwen2-Audio (7B params). With equivalent compute, a larger decoder would likely close much of this gap.\n\n4. **Temporal non-stationarity:** As shown in Novelty 2, the model degrades severely on future recording sessions (WER 95.4%). Periodic retraining or domain adaptation is needed for clinical deployment.\n\n5. **Character-level CTC:** Using a word-level language model decoder (like the BIT paper's Qwen2-Audio) would likely reduce substitution errors significantly by incorporating language priors.\n\n---\n\n## 11. Conclusion\n\nWe implemented a complete brain-to-text decoding system achieving 40.68% WER on the BT'25 validation set — a 20.5 percentage point improvement over the RNN baseline. The system combines an SSL-pretrained transformer encoder (trained using masked patch modelling) with a BiGRU CTC decoder.\n\nKey findings:\n- SSL pretraining is essential (not optional) for this task — without it, WER is 97.4%\n- Lower SSL mask ratios (0.25) outperform the standard 75% for intracortical neural data\n- Neural signals are temporally non-stationary — temporal generalisation is a critical unsolved challenge for clinical BCI deployment\n- Domain mismatch prevents direct use of audio-pretrained decoders (Whisper) — domain-specific decoders trained from scratch are necessary given current compute constraints\n\nThe implementation demonstrates that transformer-based SSL pretraining provides a strong foundation for neural speech decoding, with clear pathways for improvement through larger decoders, cross-subject training, and online adaptation methods.\n\n---\n\n","metadata":{}}]}