{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Installs, imports and global configuration","metadata":{}},{"cell_type":"code","source":"# --- 1. INSTALL LIBRARIES ---\n!pip install --quiet h5py einops jiwer\n!pip install --quiet https://github.com/kpu/kenlm/archive/master.zip pyctcdecode\n\n# --- 2. IMPORTS ---\nimport os\nimport sys\nimport re\nimport math\nimport shutil\nimport subprocess\nimport collections\nimport multiprocessing\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport h5py\nfrom tqdm.auto import tqdm\nimport jiwer\nimport time\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import OneCycleLR\nfrom torch.amp import autocast, GradScaler\n\nfrom pyctcdecode import build_ctcdecoder\n\n# --- 3. GLOBAL CONFIGURATION ---\nprint(\"Torch\", torch.__version__, \"CUDA available:\", torch.cuda.is_available())\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# This mapping is critical for decoding and comes from the competition's data description.\nLOGIT_TO_PHONEME = [\n    'BLANK', '<pad>', 'AA', 'AE', 'AH', 'AO', 'AW', 'AY', 'B', 'CH', 'D', 'DH',\n    'EH', 'ER', 'EY', 'F', 'G', 'HH', 'IH', 'IY', 'JH', 'K', 'L', 'M', 'N',\n    'NG', 'OW', 'OY', 'P', 'R', 'S', 'SH', 'T', 'TH', 'UH', 'UW', 'V', 'W',\n    'Y', 'Z', 'ZH', 'SIL'\n]","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-10T13:39:22.833352Z","iopub.execute_input":"2025-10-10T13:39:22.833651Z","iopub.status.idle":"2025-10-10T13:39:28.125647Z","shell.execute_reply.started":"2025-10-10T13:39:22.833629Z","shell.execute_reply":"2025-10-10T13:39:28.124719Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Data Preparation - Dataset, Collation and Filtering","metadata":{}},{"cell_type":"code","source":"class BrainDataset(Dataset):\n    def __init__(self, index_df, cache_size=4):\n        self.df = index_df\n        self._cache_size = cache_size\n        self._file_cache = collections.OrderedDict()\n\n    def __len__(self):\n        return len(self.df)\n\n    def _open_file(self, path):\n        if path in self._file_cache:\n            self._file_cache.move_to_end(path)\n            return self._file_cache[path]\n        f = h5py.File(path, 'r')\n        self._file_cache[path] = f\n        if len(self._file_cache) > self._cache_size:\n            old_path, old_f = self._file_cache.popitem(last=False)\n            try: old_f.close()\n            except: pass\n        return f\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        f = self._open_file(row['h5_path'])\n        g = f[row['group']]\n        feats = g['input_features'][()].astype('float32')\n        tgt = g['seq_class_ids'][()].astype('int64')\n        return torch.from_numpy(feats), torch.from_numpy(tgt)\n\ndef collate_for_ctc(batch):\n    xs, ys = zip(*batch)\n    x_lens = [x.shape[0] for x in xs]\n    t_lens = [y.shape[0] for y in ys]\n    max_t = max(x_lens)\n    channels = xs[0].shape[1]\n    x_padded = torch.zeros(len(xs), channels, max_t, dtype=torch.float32)\n    for i, x in enumerate(xs):\n        x_padded[i, :, :x.shape[0]] = x.permute(1,0)\n    targets_concat = torch.cat(ys) if sum(t_lens) > 0 else torch.tensor([], dtype=torch.long)\n    return x_padded, targets_concat, torch.tensor(x_lens, dtype=torch.long), torch.tensor(t_lens, dtype=torch.long)\n\ndef filter_dataframe_by_length(df, downsample_factor=4):\n    print(f\"Original dataframe size: {len(df)}\")\n    valid_indices = []\n    for idx, row in tqdm(df.iterrows(), total=len(df), desc=\"Filtering samples\"):\n        try:\n            with h5py.File(row['h5_path'], 'r') as hf:\n                g = hf[row['group']]\n                target_len = g['seq_class_ids'].shape[0]\n                # --- THIS IS THE FIX for the \"Zero Loss\" problem ---\n                # Ensure we only use samples with actual target text.\n                if target_len > 0 and target_len <= g['input_features'].shape[0] // downsample_factor:\n                    valid_indices.append(idx)\n        except Exception:\n            pass\n    filtered_df = df.loc[valid_indices].reset_index(drop=True)\n    print(f\"Filtered dataframe size: {len(filtered_df)} ({len(df) - len(filtered_df)} samples removed)\")\n    return filtered_df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Create Dataloaders","metadata":{}},{"cell_type":"code","source":"INPUT_DIR = Path('/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final')\n\nprint(\"Building and filtering training file index...\")\ntrain_rows = []\nfor p in INPUT_DIR.rglob(\"data_train.hdf5\"):\n    try:\n        with h5py.File(p, 'r') as hf:\n            for g in hf.keys():\n                if g.startswith('trial_'):\n                    train_rows.append({'h5_path': str(p), 'group': g})\n    except (IOError, OSError):\n        print(f\"Warning: Could not read {p}, skipping.\")\ntrain_df_unfiltered = pd.DataFrame(train_rows)\n# --- THIS IS THE FIX ---\n# Change the downsample factor to 2 to match our new model architecture.\ntrain_df_index = filter_dataframe_by_length(train_df_unfiltered, downsample_factor=2)\n# --------------------\n\nprint(\"\\nBuilding and filtering validation file index...\")\nval_rows = []\nfor p in INPUT_DIR.rglob(\"data_val.hdf5\"):\n    try:\n        with h5py.File(p, 'r') as hf:\n            for g in hf.keys():\n                if g.startswith('trial_'):\n                    val_rows.append({'h5_path': str(p), 'group': g})\n    except (IOError, OSError):\n        print(f\"Warning: Could not read {p}, skipping.\")\nval_df_unfiltered = pd.DataFrame(val_rows)\n# --- THIS IS THE FIX ---\nval_df_index = filter_dataframe_by_length(val_df_unfiltered, downsample_factor=2)\n# --------------------\n\ntrain_ds = BrainDataset(train_df_index)\nval_ds = BrainDataset(val_df_index)\n\n# Set num_workers=0 to disable multiprocessing for stability in Kaggle.\ntrain_loader = DataLoader(train_ds, batch_size=16, shuffle=True, collate_fn=collate_for_ctc, num_workers=0)\nval_loader = DataLoader(val_ds, batch_size=16, shuffle=False, collate_fn=collate_for_ctc, num_workers=0)\n\nprint(f\"\\nDataLoaders created. Train batches: {len(train_loader)}, Val batches: {len(val_loader)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Model Architecture","metadata":{}},{"cell_type":"code","source":"class PositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_len=2000):\n        super().__init__()\n        position = torch.arange(max_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))\n        pe = torch.zeros(1, max_len, d_model)\n        pe[0, :, 0::2] = torch.sin(position * div_term)\n        pe[0, :, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe)\n    def forward(self, x): return x + self.pe[:, :x.size(1)]\n\nclass ConvStem(nn.Module):\n    def __init__(self, in_ch, d_model):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv1d(in_ch, d_model // 2, kernel_size=7, stride=2, padding=3, bias=False),\n            nn.BatchNorm1d(d_model // 2), nn.GELU(),\n            # --- THIS IS THE FIX ---\n            # Reduce the stride to 1 to decrease downsampling from 4x to 2x.\n            nn.Conv1d(d_model // 2, d_model, kernel_size=5, stride=1, padding=2, bias=False),\n            # --------------------\n            nn.BatchNorm1d(d_model), nn.GELU())\n    def forward(self, x): return self.net(x)\n\nclass BrainToTextModel(nn.Module):\n    def __init__(self, in_ch=512, d_model=256, nhead=8, num_layers=6, vocab_size=len(LOGIT_TO_PHONEME)):\n        super().__init__()\n        self.conv = ConvStem(in_ch, d_model)\n        self.pos_enc = PositionalEncoding(d_model)\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead, dim_feedforward=d_model * 4,\n            dropout=0.2, activation='gelu', batch_first=True, norm_first=True)\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n        self.fc = nn.Linear(d_model, vocab_size)\n    def forward(self, x):\n        x = self.conv(x).permute(0, 2, 1)\n        x = self.pos_enc(x)\n        x = self.transformer(x)\n        logits = self.fc(x)\n        return F.log_softmax(logits, dim=-1)\n\nmodel = BrainToTextModel().to(device)\nprint(f\"Model Initialized. Total parameters: {sum(p.numel() for p in model.parameters())/1e6:.2f} Million\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Training Loop","metadata":{}},{"cell_type":"code","source":"NUM_EPOCHS = 30\nMAX_LR = 3e-4\nWEIGHT_DECAY = 0.05\nGRAD_CLIP_VALUE = 2.0\n\nctc_loss = nn.CTCLoss(blank=0, zero_infinity=True)\noptimizer = AdamW(model.parameters(), lr=MAX_LR, weight_decay=WEIGHT_DECAY)\nscaler = GradScaler()\nscheduler = OneCycleLR(optimizer, max_lr=MAX_LR, steps_per_epoch=len(train_loader), epochs=NUM_EPOCHS)\nbest_wer = float('inf')\n\nprint(f\"🚀 Starting training for {NUM_EPOCHS} epochs with OneCycleLR scheduler.\")\nfor epoch in range(1, NUM_EPOCHS + 1):\n    model.train()\n    total_loss = 0.0\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch} [train]\", leave=False)\n    for i, (x, targets, x_lens, t_lens) in enumerate(pbar):\n        if targets.numel() == 0: continue\n        \n        x, targets, x_lens, t_lens = x.to(device), targets.to(device), x_lens.to(device), t_lens.to(device)\n        \n        with autocast(device_type=device.type, dtype=torch.float16):\n            log_probs = model(x).permute(1, 0, 2)\n            # --- THIS IS THE FIX ---\n            # Divide by 2 to match the new model's downsampling factor.\n            input_lengths = torch.div(x_lens, 2, rounding_mode='floor')\n            # --------------------\n            loss = ctc_loss(log_probs, targets, input_lengths, t_lens)\n        \n        if torch.isinf(loss) or torch.isnan(loss):\n            scaler.update(); optimizer.zero_grad(); continue\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP_VALUE)\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad()\n        scheduler.step()\n            \n        total_loss += loss.item()\n        pbar.set_postfix({'loss': total_loss / (i + 1), 'lr': scheduler.get_last_lr()[0]})\n    \n    avg_train_loss = total_loss / len(train_loader) if len(train_loader) > 0 else 0\n    \n    model.eval()\n    total_val_loss, all_preds, all_refs = 0.0, [], []\n    with torch.no_grad():\n        for x, targets, x_lens, t_lens in tqdm(val_loader, desc=f\"Epoch {epoch} [val]\", leave=False):\n            if targets.numel() == 0: continue\n            \n            x, targets, x_lens, t_lens = x.to(device), targets.to(device), x_lens.to(device), t_lens.to(device)\n            log_probs = model(x)\n            # --- THIS IS THE FIX ---\n            input_lengths = torch.div(x_lens, 2, rounding_mode='floor')\n            # --------------------\n            total_val_loss += ctc_loss(log_probs.permute(1,0,2), targets, input_lengths, t_lens).item()\n            \n            decoded = log_probs.argmax(-1).cpu().numpy()\n            target_offset = 0\n            for i, L in enumerate(t_lens.cpu().numpy()):\n                pred_indices = [p for j,p in enumerate(decoded[i][:input_lengths[i]]) if (j==0 or p!=decoded[i][j-1]) and p!=0]\n                all_preds.append(\" \".join(map(str, pred_indices)))\n                all_refs.append(\" \".join(map(str, targets[target_offset:target_offset+L].cpu().numpy())))\n                target_offset += L\n    \n    avg_val_loss = total_val_loss / len(val_loader) if len(val_loader) > 0 else 0\n    wer = jiwer.wer(all_refs, all_preds) if len(all_refs) > 0 else 1.0\n    print(f\"Epoch {epoch}: Train Loss={avg_train_loss:.4f} | Val Loss={avg_val_loss:.4f} | WER={wer:.4f}\")\n\n    if wer < best_wer:\n        best_wer = wer\n        torch.save(model.state_dict(), \"/kaggle/working/best_model.pth\")\n        print(f\"✅ New best model saved with WER: {best_wer:.4f}\")\n\nprint(\"\\n✅ Training complete.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Inference and Beam Search Decoder","metadata":{}},{"cell_type":"code","source":"# === STEP 1: Build a Language Model (with integrated KenLM compilation, robust) ===\ndef build_lm(df):\n    print(\"Building Language Model from training corpus...\")\n    corpus_path = \"/kaggle/working/text_corpus.txt\"\n\n    # Collect corpus with deduplication to avoid KenLM discount errors on small/artificial data\n    unique_lines = set()\n    for idx, row in df.iterrows():\n        with h5py.File(row['h5_path'], 'r') as hf:\n            if row['group'] in hf:\n                ids = hf[row['group']]['seq_class_ids'][()]\n                phonemes = [LOGIT_TO_PHONEME[i] for i in ids]\n                line = \" \".join(phonemes)\n                if line:  # skip empty\n                    unique_lines.add(line)\n\n    with open(corpus_path, \"w\") as f:\n        for line in unique_lines:\n            f.write(line + \"\\n\")\n\n    lmplz_path = \"/kaggle/working/kenlm/build/bin/lmplz\"\n\n    if not os.path.exists(lmplz_path):\n        print(\"kenlm's lmplz not found. Cloning and building from source...\")\n        print(\"This will take a minute or two, but only happens once.\")\n        # Clone and build kenlm\n        get_ipython().system(\"git clone https://github.com/kpu/kenlm.git /kaggle/working/kenlm\")\n        get_ipython().system(\"mkdir -p /kaggle/working/kenlm/build\")\n        get_ipython().system(\"cd /kaggle/working/kenlm/build && cmake .. && make -j2\")\n        assert os.path.exists(lmplz_path), \"Build failed! lmplz executable not found after compilation.\"\n        print(\"✅ kenlm built successfully.\")\n\n    print(f\"Using lmplz at: {lmplz_path}\")\n    lm_path = \"/kaggle/working/5gram.arpa\"\n\n    # Add --discount_fallback to handle sparse counts; keep order 5 as intended\n    command = f\"{lmplz_path} --discount_fallback -o 5 < {corpus_path} > {lm_path}\"\n\n    try:\n        result = subprocess.run(command, shell=True, check=True, text=True, capture_output=True)\n    except subprocess.CalledProcessError as e:\n        print(\"--- KenLM Build Failed ---\")\n        print(\"STDOUT:\", e.stdout)\n        print(\"STDERR:\", e.stderr)\n        # As a fallback, try a lower order 4\n        print(\"Retrying with 4-gram and discount_fallback...\")\n        lm_path = \"/kaggle/working/4gram.arpa\"\n        command2 = f\"{lmplz_path} --discount_fallback -o 4 < {corpus_path} > {lm_path}\"\n        subprocess.run(command2, shell=True, check=True, text=True)\n        print(\"✅ kenlm built successfully with 4-gram fallback.\")\n        return lm_path\n\n    return lm_path\n\n# Load the best model checkpoint before building the LM\nmodel.load_state_dict(torch.load(\"/kaggle/working/best_model.pth\", map_location=device))\nlm_path = build_lm(train_df_index)\n\n# === STEP 2: Setup the CTC Beam Search Decoder ===\ndecoder = build_ctcdecoder(\n    labels=LOGIT_TO_PHONEME,\n    kenlm_model_path=lm_path,\n    alpha=0.6,\n    beta=1.0,\n)\nprint(\"✅ CTC Beam Search Decoder with LM is ready.\")\n\n# === STEP 3: Run Inference and Create Submission (chronological order, id 0..1449) ===\n\nimport re\nfrom pathlib import Path\n\ndef _parse_session_date(session_name: str):\n    # session folder like 't15.2023.08.13' -> (2023, 8, 13)\n    m = re.search(r'(\\d{4})\\.(\\d{2})\\.(\\d{2})', session_name)\n    return tuple(map(int, m.groups())) if m else (9999, 99, 99)\n\ndef _discover_test_files_in_order(input_dir: Path):\n    # Find all data_test.hdf5 and sort by session date ascending\n    items = []\n    for p in input_dir.rglob(\"data_test.hdf5\"):\n        session = p.parent.name  # e.g., 't15.2023.08.13'\n        items.append((p, session))\n    items.sort(key=lambda x: _parse_session_date(x[1]))\n    return [p for p, _ in items]\n\ndef _sorted_trials(hf):\n    trials = [k for k in hf.keys() if k.startswith(\"trial_\")]\n    def idx(k):\n        m = re.search(r'trial_(\\d+)', k)\n        return int(m.group(1)) if m else 0\n    return sorted(trials, key=idx)\n\ndef _decode_one(feats: np.ndarray) -> str:\n    # Run model and decoder; keep apostrophes, remove punctuation .,!?;:\n    with torch.no_grad():\n        x = torch.from_numpy(feats.astype(\"float32\")).unsqueeze(0).permute(0, 2, 1).to(device)\n        if device.type == \"cuda\":\n            with autocast(device_type=device.type, dtype=torch.float16):\n                log_probs = model(x)\n        else:\n            log_probs = model(x)\n    lp = log_probs.squeeze(0).detach().cpu().numpy()\n    decoded = decoder.decode(lp)\n    text = decoded[0] if isinstance(decoded, (list, tuple)) and len(decoded) > 0 else str(decoded)\n    text = re.sub(r\"[.,!?;:]\", \"\", text).strip()\n    return text\n\nprint(\"Building submission in required chronological order...\")\nordered_texts = []\nmodel.eval()\ntest_files = _discover_test_files_in_order(INPUT_DIR)\n\nfor test_path in tqdm(test_files, desc=\"Inference with LM (chronological)\"):\n    with h5py.File(test_path, \"r\") as hf:\n        for trial_name in _sorted_trials(hf):\n            feats = hf[trial_name][\"input_features\"][()].astype(\"float32\")\n            ordered_texts.append(_decode_one(feats))\n\n# Must be exactly 1450 sentences\nassert len(ordered_texts) == 1450, f\"Expected 1450 decoded sentences, got {len(ordered_texts)}\"\n\nsub_df = pd.DataFrame({\n    \"id\": list(range(1450)),\n    \"text\": ordered_texts\n})\n\n# Write to current directory as required\nsub_df.to_csv(\"submission.csv\", index=False)\n\nprint(\"\\n✅ Final submission.csv created at:\", os.path.abspath(\"submission.csv\"))\nprint(sub_df.head().to_string(index=False))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Verify format and count\ndf = pd.read_csv(\"submission.csv\")\nassert list(df.columns) == [\"id\", \"text\"], f\"Columns must be ['id','text'], got {df.columns.tolist()}\"\nassert len(df) == 1450, f\"Expected 1450 rows, got {len(df)}\"\nassert df[\"id\"].iloc[0] == 0 and df[\"id\"].iloc[-1] == 1449, \"IDs must run from 0 to 1449 in order\"\nprint(\"✅ submission.csv format & count verified.\")\nprint(df.head().to_string(index=False))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}