{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:35:09.345723Z","iopub.execute_input":"2025-12-15T05:35:09.346864Z","iopub.status.idle":"2025-12-15T05:35:11.258723Z","shell.execute_reply.started":"2025-12-15T05:35:09.346836Z","shell.execute_reply":"2025-12-15T05:35:11.257310Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.system('pip install jiwer') # Instalacja automatyczna wewnątrz skryptu\n\nimport h5py\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport json\nimport warnings\n\n# Suppress warnings for cleaner output\nwarnings.filterwarnings('ignore')\n\n# Configuration\nCONFIG = {\n    'data_dir': '/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/',\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'batch_size': 64,       # ZWIĘKSZONO z 32 na 64 (szybszy trening)\n    'num_epochs': 50,\n    'hidden_dim': 512,\n    'num_layers': 3,\n    'dropout': 0.5,         # ZWIĘKSZONO z 0.3 na 0.5 (lepsza generalizacja)\n    'learning_rate': 1e-3,\n}\n\nprint(f\"Device: {CONFIG['device']}\")\nprint(f\"PyTorch version: {torch.__version__}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:35:11.260049Z","iopub.execute_input":"2025-12-15T05:35:11.260383Z","iopub.status.idle":"2025-12-15T05:35:21.125735Z","shell.execute_reply.started":"2025-12-15T05:35:11.260363Z","shell.execute_reply":"2025-12-15T05:35:21.125049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# CTC beam search (N-best)\n# -------------------------\nimport math\nfrom collections import defaultdict, namedtuple\n\nCandidate = namedtuple('Candidate', ['seq', 'logp'])\n\ndef logsumexp(a, b):\n    if a < b:\n        a, b = b, a\n    return a + math.log1p(math.exp(b - a))\n\ndef ctc_beam_search_nbest(log_probs_tensor, beam_width=40, topk=8, blank=0):\n    \"\"\"\n    Prefix beam search in log-space returning topk candidates.\n    log_probs_tensor: torch.Tensor (T, C) of log-probs (log_softmax output)\n    Returns: list of Candidate(seq=[token ids], logp=acoustic_logprob)\n    \"\"\"\n    T, C = log_probs_tensor.shape\n    lp = log_probs_tensor.cpu().numpy()\n    prefixes = {(): (0.0, -math.inf)}  # mapping prefix -> (p_blank, p_nonblank) in log-space\n    for t in range(T):\n        new_prefixes = defaultdict(lambda: (-math.inf, -math.inf))\n        # optionally restrict expansion to top tokens at this frame for speed\n        top_indices = np.argsort(lp[t])[-min(C, beam_width):]\n        for prefix, (p_b, p_nb) in prefixes.items():\n            # extend with blank\n            pb = lp[t, blank]\n            nb0, nb1 = new_prefixes[prefix]\n            new_b = logsumexp(nb0, p_b + pb)\n            new_b = logsumexp(new_b, p_nb + pb)\n            new_prefixes[prefix] = (new_b, nb1)\n            # extend with non-blank tokens\n            for c in top_indices:\n                if c == blank:\n                    continue\n                p_c = lp[t, c]\n                new_seq = prefix + (int(c),)\n                prev_b, prev_nb = new_prefixes[new_seq]\n                if len(prefix) > 0 and prefix[-1] == c:\n                    # repetition: only extend from p_nb\n                    new_nb = logsumexp(prev_nb, p_nb + p_c)\n                    new_prefixes[new_seq] = (prev_b, new_nb)\n                else:\n                    # extend from both p_b and p_nb\n                    val = -math.inf\n                    val = logsumexp(val, p_b + p_c)\n                    val = logsumexp(val, p_nb + p_c)\n                    new_prefixes[new_seq] = (prev_b, logsumexp(prev_nb, val))\n        # prune prefixes to beam_width by total prob\n        scored = []\n        for pref, (p_b, p_nb) in new_prefixes.items():\n            total = logsumexp(p_b, p_nb)\n            scored.append((total, pref, (p_b, p_nb)))\n        scored.sort(key=lambda x: x[0], reverse=True)\n        prefixes = {pref: scores for (_, pref, scores) in scored[:beam_width]}\n    # collect final candidates\n    final = []\n    for pref, (p_b, p_nb) in prefixes.items():\n        total = logsumexp(p_b, p_nb)\n        final.append(Candidate(seq=list(pref), logp=total))\n    final.sort(key=lambda x: x.logp, reverse=True)\n    return final[:topk]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:35:21.126521Z","iopub.execute_input":"2025-12-15T05:35:21.126941Z","iopub.status.idle":"2025-12-15T05:35:21.229109Z","shell.execute_reply.started":"2025-12-15T05:35:21.126920Z","shell.execute_reply":"2025-12-15T05:35:21.228506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_split(data_dir, split='train'):\n    \"\"\"Load train/val/test split\"\"\"\n    from glob import glob\n    \n    pattern = f'{data_dir}/**/data_{split}.hdf5'\n    files = sorted(glob(pattern, recursive=True))\n    \n    print(f\"\\nLoading {split} split...\")\n    \n    all_data = {k: [] for k in ['neural', 'n_steps', 'sentence']}\n    \n    for filepath in tqdm(files):\n        with h5py.File(filepath, 'r') as f:\n            for trial_key in f.keys():\n                trial = f[trial_key]\n                \n                # Neural data\n                neural = trial['input_features'][:]\n                n_steps = trial.attrs['n_time_steps']\n                \n                # Labels (train/val only)\n                sentence = trial.attrs.get('sentence_label')\n                if sentence and isinstance(sentence, bytes):\n                    sentence = sentence.decode('utf-8')\n                \n                all_data['neural'].append(neural)\n                all_data['n_steps'].append(n_steps)\n                all_data['sentence'].append(sentence)\n    \n    print(f\"✓ Loaded {len(all_data['neural'])} samples\")\n    return all_data\n\nclass BrainToTextDataset(Dataset):\n    def __init__(self, data, char2idx=None, normalize=True):\n        self.neural = data['neural']\n        self.n_steps = data['n_steps']\n        self.sentences = data['sentence']\n        self.normalize = normalize\n        \n        if char2idx is None:\n            self.char2idx = self._build_vocab()\n        else:\n            self.char2idx = char2idx\n        \n        self.idx2char = {v: k for k, v in self.char2idx.items()}\n        self.vocab_size = len(self.char2idx)\n    \n    def _build_vocab(self):\n        chars = set()\n        for sent in self.sentences:\n            if sent:\n                chars.update(sent.lower())\n        chars = sorted(list(chars))\n        char2idx = {'<BLANK>': 0}\n        for i, ch in enumerate(chars, start=1):\n            char2idx[ch] = i\n        return char2idx\n    \n    def __len__(self):\n        return len(self.neural)\n    \n    def __getitem__(self, idx):\n        neural = self.neural[idx][:self.n_steps[idx]]\n        \n        if self.normalize:\n            neural = (neural - neural.mean()) / (neural.std() + 1e-8)\n        \n        sentence = self.sentences[idx] if self.sentences[idx] else \"\"\n        target = [self.char2idx.get(ch.lower(), 0) for ch in sentence]\n        \n        return {\n            'neural': torch.FloatTensor(neural),\n            'target': torch.LongTensor(target),\n            'length': len(neural),\n            'target_length': len(target),\n            'sentence': sentence\n        }\n\ndef collate_fn(batch):\n    batch = sorted(batch, key=lambda x: x['length'], reverse=True)\n    neurals = [item['neural'] for item in batch]\n    targets = [item['target'] for item in batch]\n    \n    neural_padded = pad_sequence(neurals, batch_first=True)\n    target_padded = pad_sequence(targets, batch_first=True)\n    lengths = torch.LongTensor([item['length'] for item in batch])\n    target_lengths = torch.LongTensor([item['target_length'] for item in batch])\n    \n    return {\n        'neural': neural_padded,\n        'target': target_padded,\n        'lengths': lengths,\n        'target_lengths': target_lengths,\n        'sentences': [item['sentence'] for item in batch]\n    }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:35:21.230012Z","iopub.execute_input":"2025-12-15T05:35:21.230277Z","iopub.status.idle":"2025-12-15T05:35:21.257525Z","shell.execute_reply.started":"2025-12-15T05:35:21.230254Z","shell.execute_reply":"2025-12-15T05:35:21.256789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GaussianNoise(nn.Module):\n    \"\"\"Adds noise to input during training to prevent overfitting\"\"\"\n    def __init__(self, sigma=0.1):\n        super().__init__()\n        self.sigma = sigma\n    \n    def forward(self, x):\n        if self.training and self.sigma > 0:\n            return x + torch.randn_like(x) * self.sigma\n        return x\n\nclass ImprovedBrainCTCModel(nn.Module):\n    def __init__(self, input_dim=512, hidden_dim=512, num_layers=3, \n                 vocab_size=50, dropout=0.5):\n        super().__init__()\n        \n        # Augmentation\n        self.noise = GaussianNoise(sigma=0.2)\n        \n        self.cnn = nn.Sequential(\n            # Layer 1: Wide kernel + Stride 2 (Time / 2)\n            nn.Conv1d(input_dim, 256, kernel_size=11, stride=2, padding=5),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            \n            # Layer 2: Standard kernel + Stride 2 (Time / 4 total)\n            nn.Conv1d(256, 256, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n        \n        self.lstm = nn.LSTM(\n            256, hidden_dim, num_layers,\n            batch_first=True, bidirectional=True,\n            dropout=dropout if num_layers > 1 else 0\n        )\n        \n        self.fc = nn.Linear(hidden_dim * 2, vocab_size)\n    \n    def forward(self, x, lengths):\n        # 1. Add noise (training only)\n        if self.training:\n            x = self.noise(x)\n            \n        # 2. CNN\n        x = x.transpose(1, 2)\n        x = self.cnn(x)\n        x = x.transpose(1, 2)\n        \n        # 3. Recalculate lengths after Stride=2 twice (div by 4)\n        cnn_lengths = (lengths.cpu() // 4).clamp(min=1)\n        \n        # 4. LSTM\n        x_packed = pack_padded_sequence(x, cnn_lengths, batch_first=True, enforce_sorted=False)\n        lstm_out, _ = self.lstm(x_packed)\n        lstm_out, _ = pad_packed_sequence(lstm_out, batch_first=True)\n        \n        # 5. Output\n        logits = self.fc(lstm_out)\n        log_probs = torch.log_softmax(logits, dim=-1)\n        \n        # Return log_probs (T, B, C) AND new lengths\n        return log_probs.transpose(0, 1), cnn_lengths\n\n# ============================================================================\n# 4. TRAINING & VALIDATION LOOPS (UPDATED)\n# ============================================================================\ndef train_model(train_loader, val_loader, char2idx, config,\n                patience=8, validate_every=1, save_path='best_model.pt'):\n    \"\"\"\n    Robust acoustic training with per-epoch validation, early stopping, checkpointing.\n    Returns: best_model (state_dict loaded into a new model instance), best_val_wer\n    \"\"\"\n    device = config['device']\n    model = ImprovedBrainCTCModel(\n        input_dim=512,\n        hidden_dim=config['hidden_dim'],\n        num_layers=config['num_layers'],\n        vocab_size=len(char2idx),\n        dropout=config['dropout']\n    ).to(device)\n\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n    optimizer = optim.AdamW(model.parameters(), lr=config['learning_rate'], weight_decay=1e-4)\n\n    # Option: use OneCycleLR if steps_per_epoch known\n    try:\n        steps = len(train_loader)\n        scheduler = optim.lr_scheduler.OneCycleLR(\n            optimizer, max_lr=config['learning_rate'], epochs=config['num_epochs'], steps_per_epoch=steps\n        )\n        use_scheduler = True\n    except Exception:\n        scheduler = None\n        use_scheduler = False\n\n    idx2char = {v: k for k, v in char2idx.items()}\n    best_wer = float('inf')\n    best_epoch = -1\n    no_improve = 0\n\n    print(\"\\nStarting acoustic training...\")\n    for epoch in range(1, config['num_epochs'] + 1):\n        model.train()\n        total_loss = 0.0\n        iters = 0\n\n        pbar = tqdm(train_loader, desc=f\"Train Epoch {epoch}/{config['num_epochs']}\", leave=False)\n        for batch in pbar:\n            neural = batch['neural'].to(device)\n            target = batch['target'].to(device)\n            lengths = batch['lengths']\n            target_lengths = batch['target_lengths']\n\n            # forward\n            log_probs, cnn_lengths = model(neural, lengths)\n            # CTC expects (T, B, C) log_probs -> already in that shape in your model\n            loss = criterion(log_probs, target, cnn_lengths, target_lengths)\n\n            optimizer.zero_grad()\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n            optimizer.step()\n            if use_scheduler:\n                scheduler.step()\n\n            total_loss += loss.item()\n            iters += 1\n            pbar.set_postfix(loss=total_loss / iters)\n\n        avg_loss = total_loss / max(1, iters)\n        print(f\"Epoch {epoch} training loss: {avg_loss:.4f}\")\n\n        # Validation\n        if epoch % validate_every == 0:\n            val_wer = validate_model(model, val_loader, idx2char, device)\n            print(f\"Epoch {epoch} VALIDATION WER: {val_wer:.2f}%\")\n\n            # checkpoint if improved\n            if val_wer < best_wer:\n                best_wer = val_wer\n                best_epoch = epoch\n                no_improve = 0\n                torch.save({\n                    'model_state_dict': model.state_dict(),\n                    'char2idx': char2idx,\n                    'config': config,\n                    'wer': best_wer,\n                    'epoch': epoch\n                }, save_path)\n                print(f\"✓ New best model saved (WER {best_wer:.2f}%) at epoch {epoch}\")\n            else:\n                no_improve += 1\n                print(f\"No improvement for {no_improve} validations (patience {patience})\")\n\n            # early stopping\n            if no_improve >= patience:\n                print(f\"Early stopping after {epoch} epochs (best epoch {best_epoch}, best WER {best_wer:.2f}%)\")\n                break\n\n    # load best into model instance to return\n    if os.path.exists(save_path):\n        ckpt = torch.load(save_path, map_location=device)\n        model.load_state_dict(ckpt['model_state_dict'])\n        print(f\"Loaded best checkpoint from epoch {ckpt.get('epoch', 'unknown')} (WER {ckpt.get('wer', 'unknown')})\")\n    else:\n        print(\"No checkpoint saved during training; returning last model\")\n\n    return model, best_wer\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:35:21.258334Z","iopub.execute_input":"2025-12-15T05:35:21.258571Z","iopub.status.idle":"2025-12-15T05:35:21.278774Z","shell.execute_reply.started":"2025-12-15T05:35:21.258549Z","shell.execute_reply":"2025-12-15T05:35:21.277699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# Tiny char-level LSTM LM\n# -------------------------\nimport torch.nn.functional as F\n\nclass TinyCharLM(nn.Module):\n    def __init__(self, vocab_size, emb=128, hidden=256, nlayers=2, dropout=0.2, pad_idx=0):\n        super().__init__()\n        self.emb = nn.Embedding(vocab_size, emb, padding_idx=pad_idx)\n        self.lstm = nn.LSTM(emb, hidden, nlayers, batch_first=True, dropout=dropout)\n        self.fc = nn.Linear(hidden, vocab_size)\n    def forward(self, x, lengths):\n        e = self.emb(x)\n        packed = pack_padded_sequence(e, lengths.cpu(), batch_first=True, enforce_sorted=False)\n        out, _ = self.lstm(packed)\n        out, _ = pad_packed_sequence(out, batch_first=True)\n        logits = self.fc(out)\n        return logits\n\ndef build_lm_corpus(train_dataset, val_dataset):\n    def get_sents(ds):\n        s = []\n        for i in range(len(ds)):\n            sent = ds.sentences[i]\n            if sent and len(sent.strip())>0:\n                s.append(sent.lower().strip())\n        return s\n    train_sents = get_sents(train_dataset)\n    val_sents = get_sents(val_dataset)\n    return train_sents, val_sents\n\nclass CharLMDataset(Dataset):\n    def __init__(self, sentences, char2idx):\n        self.data = []\n        for s in sentences:\n            idxs = [char2idx.get(ch, 0) for ch in s]\n            if len(idxs) < 2: continue\n            self.data.append((torch.LongTensor(idxs[:-1]), torch.LongTensor(idxs[1:])))\n    def __len__(self): return len(self.data)\n    def __getitem__(self, idx): return self.data[idx]\n\ndef collate_lm(batch):\n    inputs = [x[0] for x in batch]\n    targets = [x[1] for x in batch]\n    inputs = pad_sequence(inputs, batch_first=True, padding_value=0)\n    targets = pad_sequence(targets, batch_first=True, padding_value=0)\n    lengths = torch.LongTensor([len(x) for x in batch])\n    return inputs, targets, lengths\n\ndef train_lm_on_corpus(train_dataset, val_dataset, char2idx, device,\n                       lm_epochs=6, lm_batch=256, lm_save='lm_best.pt', patience=3):\n    \"\"\"\n    Train small char LM with epochs and early stopping. Returns trained lm_model or None.\n    \"\"\"\n    train_sents, val_sents = build_lm_corpus(train_dataset, val_dataset)\n    if len(train_sents) == 0:\n        print(\"No LM training sentences found.\")\n        return None\n\n    train_ds = CharLMDataset(train_sents, char2idx)\n    val_ds = CharLMDataset(val_sents, char2idx)\n    train_loader = DataLoader(train_ds, batch_size=lm_batch, shuffle=True, collate_fn=collate_lm)\n    val_loader = DataLoader(val_ds, batch_size=lm_batch, shuffle=False, collate_fn=collate_lm)\n\n    lm_model = TinyCharLM(vocab_size=len(char2idx), emb=128, hidden=256, nlayers=2, dropout=0.2, pad_idx=0).to(device)\n    opt = optim.AdamW(lm_model.parameters(), lr=3e-3, weight_decay=1e-5)\n    crit = nn.CrossEntropyLoss(ignore_index=0)\n\n    best_val_loss = float('inf')\n    no_improve = 0\n\n    for ep in range(1, lm_epochs + 1):\n        lm_model.train()\n        total = 0.0\n        for inputs, targets, lengths in tqdm(train_loader, desc=f\"LM Train {ep}/{lm_epochs}\", leave=False):\n            inputs = inputs.to(device); targets = targets.to(device); lengths = lengths.to(device)\n            logits = lm_model(inputs, lengths)\n            L = logits.view(-1, logits.size(-1))\n            T = targets.view(-1)\n            loss = crit(L, T)\n            opt.zero_grad(); loss.backward()\n            torch.nn.utils.clip_grad_norm_(lm_model.parameters(), 1.0)\n            opt.step()\n            total += loss.item()\n        avg_train = total / max(1, len(train_loader))\n        # validate\n        lm_model.eval()\n        vtotal = 0.0\n        with torch.no_grad():\n            for inputs, targets, lengths in val_loader:\n                inputs = inputs.to(device); targets = targets.to(device); lengths = lengths.to(device)\n                logits = lm_model(inputs, lengths)\n                L = logits.view(-1, logits.size(-1))\n                T = targets.view(-1)\n                loss = crit(L, T)\n                vtotal += loss.item()\n        avg_val = vtotal / max(1, len(val_loader))\n        print(f\"LM Epoch {ep}/{lm_epochs} train_loss={avg_train:.4f} val_loss={avg_val:.4f}\")\n\n        if avg_val < best_val_loss:\n            best_val_loss = avg_val\n            no_improve = 0\n            torch.save({'state': lm_model.state_dict(), 'char2idx': char2idx}, lm_save)\n            print(f\"✓ LM improved and saved (val_loss {best_val_loss:.4f})\")\n        else:\n            no_improve += 1\n            print(f\"LM no improvement {no_improve}/{patience}\")\n        if no_improve >= patience:\n            print(\"LM early stopping\")\n            break\n\n    # load best LM\n    if os.path.exists(lm_save):\n        ck = torch.load(lm_save, map_location=device)\n        lm_model.load_state_dict(ck['state'])\n        return lm_model\n    return lm_model  # return last model if no saved ckpt\n\n\ndef score_sequence_lm(seq_ids, lm_model, device):\n    \"\"\"Return log-prob (natural log) of seq_ids under LM. seq_ids includes full sequence tokens.\"\"\"\n    if lm_model is None or len(seq_ids) < 2:\n        return -1e9\n    lm_model.eval()\n    with torch.no_grad():\n        inp = torch.LongTensor(seq_ids[:-1]).unsqueeze(0).to(device)\n        lengths = torch.LongTensor([inp.size(1)])\n        logits = lm_model(inp, lengths)  # (1, L-1, V)\n        logp = F.log_softmax(logits, dim=-1)\n        targets = torch.LongTensor(seq_ids[1:]).unsqueeze(0).to(device)\n        gathered = logp.gather(2, targets.unsqueeze(-1)).squeeze(-1)\n        total_logp = gathered.sum().item()\n    return total_logp\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:35:21.280326Z","iopub.execute_input":"2025-12-15T05:35:21.280590Z","iopub.status.idle":"2025-12-15T05:35:21.299918Z","shell.execute_reply.started":"2025-12-15T05:35:21.280573Z","shell.execute_reply":"2025-12-15T05:35:21.299201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# Rescoring utilities\n# -------------------------\ndef tokens_to_text(token_ids, idx2char):\n    s = []\n    prev = None\n    for t in token_ids:\n        if t == 0:\n            prev = t\n            continue\n        if t != prev:\n            s.append(idx2char.get(t, ''))\n        prev = t\n    return ''.join(s)\n\ndef rescore_and_choose(nbest_candidates, lm_model, char2idx, idx2char, alpha=1.0, beta=0.5, length_norm=True):\n    \"\"\"\n    nbest_candidates: list of Candidate(seq=[ids...], logp=acoustic_logprob)\n    alpha: acoustic weight, beta: LM weight\n    length_norm: divide acoustic logprob by T (and LM by L) to normalize lengths\n    \"\"\"\n    best_score = -1e12\n    best_text = \"\"\n    for cand in nbest_candidates:\n        seq = cand.seq\n        text = tokens_to_text(seq, idx2char)\n        if len(text) == 0:\n            lm_logp = -1e9\n        else:\n            # LM expects token ids (no blanks) -> map chars to ids via char2idx\n            seq_ids = [char2idx.get(ch, 0) for ch in text]\n            # optionally add start token if you used one; here we keep plain\n            lm_logp = score_sequence_lm([0] + seq_ids, lm_model, CONFIG['device'])\n        acoustic = cand.logp\n        if length_norm:\n            acoustic = acoustic / max(1, len(seq))\n            lm_logp = lm_logp / max(1, len(seq_ids)) if len(seq_ids)>0 else lm_logp\n        total = alpha * acoustic + beta * lm_logp\n        if total > best_score:\n            best_score = total\n            best_text = text\n    return best_text, best_score\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:35:21.301453Z","iopub.execute_input":"2025-12-15T05:35:21.301821Z","iopub.status.idle":"2025-12-15T05:35:21.319313Z","shell.execute_reply.started":"2025-12-15T05:35:21.301803Z","shell.execute_reply":"2025-12-15T05:35:21.318537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------\n# Generate predictions using beam search + LM rescoring\n# -------------------------\ndef generate_predictions_with_rescoring(model, test_samples, char2idx, idx2char, lm_model, device,\n                                       beam_width=40, topk=8, alpha=1.0, beta=0.5, batch_size=16):\n    model.eval()\n    predictions = []\n    print(\"\\nGenerating predictions with beam+LM rescoring...\")\n    with torch.no_grad():\n        for i in tqdm(range(0, len(test_samples), batch_size)):\n            batch = test_samples[i:i+batch_size]\n            features = [s['features'] for s in batch]\n            lengths = torch.LongTensor([len(f) for f in features])\n            feats_padded = pad_sequence(features, batch_first=True).to(device)\n            # sort\n            sorted_lengths, sorted_idx = lengths.sort(descending=True)\n            feats_sorted = feats_padded[sorted_idx]\n            log_probs, cnn_lengths = model(feats_sorted, sorted_lengths)  # (T, B, C)\n            # for each sample in batch:\n            batch_preds = []\n            for b in range(log_probs.size(1)):\n                T_b = cnn_lengths[b].item()\n                lp = log_probs[:T_b, b, :].detach()\n                nbest = ctc_beam_search_nbest(lp, beam_width=beam_width, topk=topk, blank=0)\n                best_text, _ = rescore_and_choose(nbest, lm_model, char2idx, idx2char, alpha=alpha, beta=beta)\n                batch_preds.append(best_text)\n            # unsort\n            unsorted = [''] * len(batch_preds)\n            for j, pred in zip(sorted_idx.tolist(), batch_preds):\n                unsorted[j] = pred\n            predictions.extend(unsorted)\n    return predictions\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:35:21.320006Z","iopub.execute_input":"2025-12-15T05:35:21.320246Z","iopub.status.idle":"2025-12-15T05:35:21.342582Z","shell.execute_reply.started":"2025-12-15T05:35:21.320228Z","shell.execute_reply":"2025-12-15T05:35:21.341889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    print(\"=\"*80)\n    print(\"BRAIN-TO-TEXT '25 - IMPROVED PIPELINE\")\n    print(\"=\"*80)\n\n    # 1. Load Data\n    train_data = load_split(CONFIG['data_dir'], 'train')\n    val_data = load_split(CONFIG['data_dir'], 'val')\n\n    # 2. Datasets\n    train_dataset = BrainToTextDataset(train_data)\n    val_dataset = BrainToTextDataset(val_data, char2idx=train_dataset.char2idx)\n\n    train_loader = DataLoader(train_dataset, batch_size=CONFIG['batch_size'],\n                              shuffle=True, collate_fn=collate_fn, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=CONFIG['batch_size'],\n                            shuffle=False, collate_fn=collate_fn, num_workers=2)\n\n    # 3. Acoustic training (explicit epochs, early stopping)\n    model, best_wer = train_model(train_loader, val_loader, train_dataset.char2idx, CONFIG,\n                                  patience=8, validate_every=1, save_path='best_model.pt')\n    print(f\"\\n✓ Acoustic training complete! Best WER: {best_wer:.2f}%\")\n\n    # 4. Load test data\n    test_samples = load_test_data_for_submission(CONFIG['data_dir'])\n\n    # 5. Load best acoustic checkpoint (already loaded in train_model return but re-load to be safe)\n    checkpoint = torch.load('best_model.pt', map_location=CONFIG['device'])\n    model.load_state_dict(checkpoint['model_state_dict'])\n    idx2char = {v: k for k, v in checkpoint['char2idx'].items()}\n\n    # 6. Train LM (epochs + validation)\n    lm_epochs = 6\n    print(f\"\\nTraining LM for {lm_epochs} epochs...\")\n    lm_model = train_lm_on_corpus(train_dataset, val_dataset, train_dataset.char2idx,\n                                  CONFIG['device'], lm_epochs=lm_epochs, lm_batch=256, lm_save='lm_best.pt')\n\n    # 7. Tune rescoring weights on validation (grid search)\n    if lm_model is None:\n        print(\"LM not available — falling back to greedy decode for test.\")\n        predictions = generate_predictions(model, test_samples, idx2char, CONFIG['device'])\n    else:\n        from jiwer import wer\n        refs = [s.lower() if s else \"\" for s in val_dataset.sentences]\n\n        candidate_alphas = [0.6, 0.8, 1.0, 1.2]\n        candidate_betas = [0.2, 0.4, 0.6]\n\n        best_cfg = (1.0, 0.5)\n        best_val_wer = 1.0\n\n        print(\"\\nGrid search rescoring weights (will use beam_width=30, topk=6 for speed)...\")\n        for a in candidate_alphas:\n            for b in candidate_betas:\n                preds = []\n                with torch.no_grad():\n                    for batch in DataLoader(val_dataset, batch_size=8, shuffle=False, collate_fn=collate_fn):\n                        neural = batch['neural'].to(CONFIG['device'])\n                        lengths = batch['lengths']\n                        log_probs, cnn_lengths = model(neural, lengths)\n                        for bb in range(log_probs.size(1)):\n                            T_b = cnn_lengths[bb].item()\n                            lp = log_probs[:T_b, bb, :].detach()\n                            nbest = ctc_beam_search_nbest(lp, beam_width=30, topk=6, blank=0)\n                            best_text, _ = rescore_and_choose(nbest, lm_model, train_dataset.char2idx, idx2char, alpha=a, beta=b)\n                            preds.append(best_text.lower())\n                try:\n                    cur_wer = wer(refs[:len(preds)], preds)\n                except:\n                    cur_wer = 1.0\n                print(f\"alpha={a} beta={b} -> WER={cur_wer*100:.2f}%\")\n                if cur_wer < best_val_wer:\n                    best_val_wer = cur_wer\n                    best_cfg = (a, b)\n        print(\"Best rescoring config:\", best_cfg, \"val WER:\", best_val_wer*100)\n\n        # 8. Generate final predictions on test using best config\n        a, b = best_cfg\n        predictions = generate_predictions_with_rescoring(model, test_samples, train_dataset.char2idx,\n                                                         idx2char, lm_model, device=CONFIG['device'],\n                                                         beam_width=40, topk=8, alpha=a, beta=b, batch_size=8)\n\n    # 9. Save submission\n    df = pd.DataFrame({'id': [s['id'] for s in test_samples], 'text': predictions})\n    df = df.sort_values('id')\n    df.to_csv('submission.csv', index=False)\n    print(\"\\n✓ submission.csv ready!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T05:36:50.899616Z","iopub.execute_input":"2025-12-15T05:36:50.900224Z","iopub.status.idle":"2025-12-15T05:36:50.911719Z","shell.execute_reply.started":"2025-12-15T05:36:50.900199Z","shell.execute_reply":"2025-12-15T05:36:50.911008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}