{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nos.system(\"pip install jiwer --quiet\")  # Quiet install\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 glob import glob\nfrom tqdm import tqdm\nfrom jiwer import wer\nimport warnings\nimport torch.nn.functional as F\nwarnings.filterwarnings(\"ignore\")\n\n# ================================================================\n# FIXED CONFIG (Proven Values)\n# ================================================================\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\": 65,      # Stable batch size\n    \"epochs\": 50,\n    \"lr\": 3e-4,            # Proven optimal LR\n    \"feature_dim\": 512,\n    \"model_dim\": 512,\n    \"num_heads\": 6,\n    \"num_layers\": 6,       # Stable 3 layers\n    \"dropout\": 0.3,\n}\n\nprint(f\" Device: {CONFIG['device']}\")\nprint(f\" PyTorch: {torch.__version__}\")\n\n# ================================================================\n#  FIXED DATA LOADING (Robust Error Handling)\n# ================================================================\ndef load_split(split):\n    \"\"\" FIXED: Handles all HDF5 edge cases\"\"\"\n    files = sorted(glob(f\"{CONFIG['data_dir']}/**/data_{split}.hdf5\", recursive=True))\n    data = {\"neural\": [], \"n_steps\": [], \"sentence\": []}\n\n    print(f\" Loading {split} data ({len(files)} files)...\")\n    \n    for fp in tqdm(files, desc=f\"Loading {split}\"):\n        try:\n            with h5py.File(fp, \"r\") as f:\n                for k in f.keys():\n                    trial = f[k]\n                    if \"input_features\" not in trial:\n                        continue\n                    \n                    # Neural data\n                    neural_raw = trial[\"input_features\"][:]\n                    n_steps = trial.attrs.get(\"n_time_steps\", len(neural_raw))\n                    \n                    #  ROBUST SENTENCE EXTRACTION\n                    sentence_attr = trial.attrs.get(\"sentence_label\", \"\")\n                    if isinstance(sentence_attr, bytes):\n                        sentence = sentence_attr.decode(\"utf-8\", errors=\"ignore\").strip()\n                    else:\n                        sentence = str(sentence_attr).strip()\n                    \n                    if len(neural_raw) == 0 or len(sentence.strip()) == 0:\n                        continue\n                    \n                    # Truncate to valid length\n                    neural = neural_raw[:n_steps]\n                    data[\"neural\"].append(neural)\n                    data[\"n_steps\"].append(n_steps)\n                    data[\"sentence\"].append(sentence)\n                    \n        except Exception as e:\n            continue  # Skip corrupt files silently\n    \n    print(f\" {split}: {len(data['neural'])} valid samples\")\n    return data\n\n# ================================================================\n# FIXED DATASET CLASS\n# ================================================================\nclass BrainDataset(Dataset):\n    def __init__(self, data, char2idx=None):\n        self.X = data[\"neural\"]\n        self.L = data[\"n_steps\"]\n        self.Y = data[\"sentence\"]\n\n        #  Build vocab safely\n        if char2idx is None:\n            all_text = \"\".join([s.lower() for s in self.Y if s])\n            chars = sorted(set(all_text))\n            self.char2idx = {\"<BLANK>\": 0, \"<UNK>\": 1}\n            for i, c in enumerate(chars, 2):\n                self.char2idx[c] = i\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        print(f\" Vocab size: {self.vocab_size}\")\n\n    def __len__(self):\n        return len(self.X)\n\n    def __getitem__(self, i):\n        # Normalize neural data\n        x = self.X[i]\n        x = (x - x.mean(axis=0)) / (x.std(axis=0) + 1e-8)\n        \n        y = self.Y[i] or \"\"\n        tgt = [self.char2idx.get(c.lower(), 1) for c in y]  # 1 = <UNK>\n        \n        return {\n            \"x\": torch.FloatTensor(x),\n            \"y\": torch.LongTensor(tgt),\n            \"xl\": len(x),\n            \"yl\": len(tgt),\n            \"sent\": y\n        }\n\ndef collate_fn(batch):\n    \"\"\" Fixed collate function\"\"\"\n    batch = sorted(batch, key=lambda x: x[\"xl\"], reverse=True)\n    \n    return {\n        \"x\": pad_sequence([b[\"x\"] for b in batch], batch_first=True),\n        \"y\": pad_sequence([b[\"y\"] for b in batch], batch_first=True),\n        \"xl\": torch.LongTensor([b[\"xl\"] for b in batch]),\n        \"yl\": torch.LongTensor([b[\"yl\"] for b in batch]),\n        \"sent\": [b[\"sent\"] for b in batch]\n    }\n\n# ================================================================\n#  STABLE MODEL: CNN + BiLSTM + CTC (Production Ready)\n# ================================================================\nclass ProductionBrainModel(nn.Module):\n    def __init__(self, input_dim=512, vocab_size=50, hidden_dim=512):\n        super().__init__()\n        \n        # CNN Feature Extractor (Time reduction 4x)\n        self.cnn = nn.Sequential(\n            nn.Conv1d(input_dim, 256, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            \n            nn.Conv1d(256, 128, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.2)\n        )\n        \n        # BiLSTM (Temporal modeling)\n        self.lstm = nn.LSTM(\n            input_size=128,\n            hidden_size=hidden_dim//2,  # //2 because bidirectional\n            num_layers=6,\n            batch_first=True,\n            bidirectional=True,\n            dropout=0.3\n        )\n        \n        # Output projection\n        self.proj = nn.Linear(hidden_dim, vocab_size)\n        \n    def forward(self, x, lengths):\n        # x: (B, T, F=512) → CNN expects (B, F, T)\n        x = x.transpose(1, 2)\n        x = self.cnn(x)  # (B, 256, T/4)\n        x = x.transpose(1, 2)  # (B, T/4, 256)\n        \n        # Update lengths after CNN (time reduced 4x)\n        cnn_lengths = (lengths // 4).clamp(min=1)\n        \n        # Pack for LSTM (variable length)\n        x_packed = pack_padded_sequence(x, cnn_lengths.cpu(), 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        # CTC logits\n        logits = self.proj(lstm_out)\n        log_probs = F.log_softmax(logits, dim=-1)\n        \n        return log_probs.transpose(0, 1), cnn_lengths  # (T, B, V)\n\n# ================================================================\n# EVALUATION FUNCTION (Fixed)\n# ================================================================\ndef evaluate(model, loader, idx2char):\n    \"\"\" Fixed evaluation with proper CTC decoding\"\"\"\n    model.eval()\n    P, T = [], []\n\n    with torch.no_grad():\n        for b in loader:\n            x = b[\"x\"].to(CONFIG[\"device\"])\n            lengths = b[\"xl\"]\n            \n            log_probs, output_lengths = model(x, lengths)\n            \n            # CTC Greedy Decoding\n            for i in range(log_probs.size(1)):\n                # Get sequence for this batch item\n                seq = log_probs[:output_lengths[i], i].argmax(-1).cpu().numpy()\n                \n                # Collapse repeats + remove blanks\n                decoded = []\n                prev_idx = None\n                for idx in seq:\n                    if idx != 0 and idx != prev_idx:  # 0 = blank\n                        decoded.append(idx2char.get(idx, \"\"))\n                    prev_idx = idx\n                \n                P.append(\"\".join(decoded))\n            \n            T.extend([s.strip() for s in b[\"sent\"]])\n\n    try:\n        w = wer(T, P)\n        return w * 100\n    except:\n        return 100.0\n\n# ================================================================\n# TRAINING LOOP (Production Ready)\n# ================================================================\ndef train(train_loader, val_loader, dataset):\n    \"\"\" Stable training loop\"\"\"\n    model = ProductionBrainModel(\n        input_dim=CONFIG[\"feature_dim\"], \n        vocab_size=len(dataset.char2idx)\n    ).to(CONFIG[\"device\"])\n    \n    optimizer = optim.AdamW(model.parameters(), lr=CONFIG[\"lr\"], weight_decay=1e-4)\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n    \n    best_wer = float('inf')\n    \n    print(f\" Training model (Vocab: {len(dataset.char2idx)})...\")\n    \n    for epoch in range(CONFIG[\"epochs\"]):\n        # Training\n        model.train()\n        total_loss = 0\n        \n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{CONFIG['epochs']}\")\n        for b in pbar:\n            x = b[\"x\"].to(CONFIG[\"device\"])\n            y = b[\"y\"].to(CONFIG[\"device\"])\n            input_lengths = b[\"xl\"]\n            target_lengths = b[\"yl\"]\n            \n            optimizer.zero_grad()\n            \n            # Forward\n            log_probs, output_lengths = model(x, input_lengths)\n            \n            # CTC Loss\n            loss = criterion(log_probs, y, output_lengths, target_lengths)\n            \n            # Backward\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            \n            total_loss += loss.item()\n            pbar.set_postfix({\"Loss\": f\"{loss.item():.4f}\"})\n        \n        avg_loss = total_loss / len(train_loader)\n        \n        # Validation\n        if (epoch + 1) % 5 == 0:\n            val_wer = evaluate(model, val_loader, dataset.idx2char)\n            print(f\"\\n Epoch {epoch+1}: Loss={avg_loss:.4f}, Val WER={val_wer:.2f}%\")\n            \n            # Save best model\n            if val_wer < best_wer:\n                best_wer = val_wer\n                torch.save({\n                    \"model\": model.state_dict(),\n                    \"char2idx\": dataset.char2idx,\n                    \"config\": CONFIG,\n                    \"wer\": val_wer\n                }, \"best_model.pt\")\n                print(f\" NEW BEST: {best_wer:.2f}% WER\")\n    \n    print(f\"\\n FINAL BEST WER: {best_wer:.2f}%\")\n    return model\n\n# ================================================================\n# TEST INFERENCE + SUBMISSION\n# ================================================================\ndef load_test_data():\n    \"\"\" Fixed test data loading\"\"\"\n    files = sorted(glob(f\"{CONFIG['data_dir']}/**/data_test.hdf5\", recursive=True))\n    samples = []\n    sid = 0\n\n    print(\" Loading test data...\")\n    for fp in tqdm(files, desc=\"Test files\"):\n        try:\n            with h5py.File(fp, \"r\") as f:\n                for k in f.keys():\n                    trial = f[k]\n                    if \"input_features\" not in trial:\n                        continue\n                    \n                    x_raw = trial[\"input_features\"][:]\n                    n_steps = trial.attrs.get(\"n_time_steps\", len(x_raw))\n                    \n                    # Normalize\n                    x = x_raw[:n_steps]\n                    x = (x - x.mean(axis=0)) / (x.std(axis=0) + 1e-8)\n                    \n                    samples.append({\"id\": sid, \"x\": torch.FloatTensor(x)})\n                    sid += 1\n        except:\n            continue\n    \n    print(f\" Loaded {len(samples)} test samples\")\n    return samples\n\ndef generate_submission(model, samples, idx2char):\n    \"\"\" Fixed inference with proper decoding\"\"\"\n    model.eval()\n    predictions = []\n    \n    print(\" Generating predictions...\")\n    with torch.no_grad():\n        for i in tqdm(range(0, len(samples), CONFIG[\"batch_size\"])):\n            batch = samples[i:i+CONFIG[\"batch_size\"]]\n            \n            # Pad batch\n            xs = [s[\"x\"] for s in batch]\n            lengths = torch.LongTensor([len(x) for x in xs])\n            x_padded = pad_sequence(xs, batch_first=True).to(CONFIG[\"device\"])\n            \n            # Sort by length (required for pack_padded_sequence)\n            sorted_lengths, sort_idx = lengths.sort(descending=True)\n            x_sorted = x_padded[sort_idx]\n            \n            # Inference\n            log_probs, output_lengths = model(x_sorted, sorted_lengths)\n            \n            # Decode (unsort afterwards)\n            for j in range(log_probs.size(1)):\n                # CTC decoding\n                seq = log_probs[:output_lengths[j], j].argmax(-1).cpu().numpy()\n                decoded = []\n                prev = None\n                for token in seq:\n                    if token != 0 and token != prev:\n                        decoded.append(idx2char.get(token, \"\"))\n                    prev = token\n                predictions.append(\"\".join(decoded))\n            \n            # Reorder predictions\n            unsorted_preds = [\"\"] * len(batch)\n            for orig_idx, sorted_idx in enumerate(sort_idx):\n                unsorted_preds[orig_idx] = predictions[sorted_idx.item()]\n            predictions[:len(batch)] = unsorted_preds[:len(batch)]\n    \n    return predictions\n\n# ================================================================\n# MAIN EXECUTION\n# ================================================================\ndef main():\n    print(\"=\" * 80)\n    print(\" BRAIN-TO-TEXT 2025 — PRODUCTION READY\")\n    print(\"=\" * 80)\n    \n    # 1. Load Data\n    train_data = load_split(\"train\")\n    val_data = load_split(\"val\")\n    \n    if len(train_data[\"neural\"]) == 0:\n        print(\" No training data found!\")\n        return\n    \n    # 2. Datasets & Loaders\n    train_ds = BrainDataset(train_data)\n    val_ds = BrainDataset(val_data, char2idx=train_ds.char2idx)\n    \n    train_dl = DataLoader(train_ds, batch_size=CONFIG[\"batch_size\"], \n                         shuffle=True, collate_fn=collate_fn, num_workers=0)\n    val_dl = DataLoader(val_ds, batch_size=CONFIG[\"batch_size\"], \n                       shuffle=False, collate_fn=collate_fn, num_workers=0)\n    \n    # 3. Train\n    model = train(train_dl, val_dl, train_ds)\n    \n    # 4. Load best model & predict test\n    ckpt = torch.load(\"best_model.pt\", map_location=CONFIG[\"device\"])\n    model = ProductionBrainModel(\n        input_dim=CONFIG[\"feature_dim\"], \n        vocab_size=len(ckpt[\"char2idx\"])\n    ).to(CONFIG[\"device\"])\n    model.load_state_dict(ckpt[\"model\"])\n    \n    idx2char = {v: k for k, v in ckpt[\"char2idx\"].items()}\n    \n    # 5. Generate submission\n    test_samples = load_test_data()\n    predictions = generate_submission(model, test_samples, idx2char)\n    \n    # 6. Save CSV\n    df = pd.DataFrame({\n        \"id\": [s[\"id\"] for s in test_samples],\n        \"text\": predictions\n    }).sort_values(\"id\").reset_index(drop=True)\n    \n    df.to_csv(\"submission.csv\", index=False)\n    print(f\"\\n SUBMISSION GENERATED: {len(df)} predictions\")\n    print(df.head())\n    print(\" READY FOR KAGGLE SUBMISSION!\")\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T13:32:33.469356Z","iopub.execute_input":"2025-12-29T13:32:33.469936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}