{"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":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nDSD-NLA: Neural Encoder → Text Decoder\nSimplified implementation (no diffusion prior, no cross-modal alignment).\n\nArchitecture:\n1. Neural Encoder (Transformer/Conformer) → z_neural\n2. Text Decoder (Transformer LM) → text (no phonemes, no CTC)\n\nEnd-to-end: neural → text\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport math\n\n# ============================================================================\n# 1. NEURAL ENCODER \n# ============================================================================\n\nclass NeuralEncoder(nn.Module):\n    \"\"\"\n    Encode neural signals (512 channels, T timesteps) → latent z_neural\n    \"\"\"\n    def __init__(self, n_channels: int = 512, d_model: int = 512,\n                 n_layers: int = 8, n_heads: int = 8, dropout: float = 0.1):\n        super().__init__()\n        \n        # Input projection - 512 channels → d_model\n        self.input_proj = nn.Sequential(\n            nn.Conv1d(n_channels, d_model, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(d_model),\n            nn.GELU(),\n            nn.Dropout(dropout)\n        )\n        \n        # Positional encoding\n        self.pos_enc = PositionalEncoding(d_model, dropout, max_len=5000)\n        \n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=n_heads,\n            dim_feedforward=d_model * 4,\n            dropout=dropout,\n            activation='gelu',\n            batch_first=True,\n            norm_first=True  # Pre-LN for stability\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=n_layers)\n        \n        self.layer_norm = nn.LayerNorm(d_model)\n        \n    def forward(self, x: torch.Tensor, mask: torch.Tensor | None = None) -> torch.Tensor:\n        \"\"\"\n        x: (B, T, 512) neural features\n        mask: (B, T//2) optional src_key_padding_mask\n        Returns: (B, L, d_model) neural latent\n        \"\"\"\n        # Conv projection + downsample 2x\n        x = x.transpose(1, 2)          # (B, 512, T)\n        x = self.input_proj(x)         # (B, d_model, T//2)\n        x = x.transpose(1, 2)          # (B, T//2, d_model)\n        \n        # Add positional encoding\n        x = self.pos_enc(x)\n        \n        # Transformer encoding\n        z_neural = self.transformer(x, src_key_padding_mask=mask)\n        z_neural = self.layer_norm(z_neural)\n        \n        return z_neural\n\n\n# ============================================================================\n# 2. TEXT DECODER - Direct neural → text (NO PHONEMES!)\n# ============================================================================\n\nclass TextDecoder(nn.Module):\n    \"\"\"\n    Decode from neural latent directly to text\n    \n    Not using phoneme intermediate representation!\n    Learn end-to-end mapping from brain → words\n    \n    Using Transformer decoder with:\n    - Autoregressive generation\n    - Character-level or BPE tokenization\n    \"\"\"\n    def __init__(self, d_model: int = 512, vocab_size: int = 256,\n                 n_layers: int = 6, n_heads: int = 8, dropout: float = 0.1):\n        super().__init__()\n        \n        self.d_model = d_model\n        self.vocab_size = vocab_size\n        \n        # Token embedding (characters or BPE)\n        self.token_embedding = nn.Embedding(vocab_size, d_model)\n        self.pos_enc = PositionalEncoding(d_model, dropout)\n        \n        # Transformer Decoder\n        decoder_layer = nn.TransformerDecoderLayer(\n            d_model=d_model,\n            nhead=n_heads,\n            dim_feedforward=d_model * 4,\n            dropout=dropout,\n            activation='gelu',\n            batch_first=True,\n            norm_first=True\n        )\n        self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=n_layers)\n        \n        # Output projection\n        self.output_proj = nn.Linear(d_model, vocab_size)\n        \n    def forward(\n        self,\n        z_neural: torch.Tensor,\n        target_tokens: torch.Tensor | None = None,\n        max_len: int = 200,\n    ) -> torch.Tensor:\n        \"\"\"\n        z_neural: (B, L, d_model) encoded neural features\n        target_tokens: (B, T) ground truth tokens (during training)\n        \n        Returns:\n            (B, T, vocab_size) logits if target_tokens is not None,\n            otherwise generated tokens via self.generate(...)\n        \"\"\"\n        if target_tokens is not None:\n            # Teacher forcing (training)\n            tgt_emb = self.token_embedding(target_tokens)\n            tgt_emb = self.pos_enc(tgt_emb)\n            \n            # Create causal mask\n            T = target_tokens.size(1)\n            causal_mask = nn.Transformer.generate_square_subsequent_mask(T).to(z_neural.device)\n            \n            # Decode\n            out = self.decoder(tgt_emb, z_neural, tgt_mask=causal_mask)\n            logits = self.output_proj(out)\n            return logits\n        else:\n            # Autoregressive generation (inference)\n            return self.generate(z_neural, max_len=max_len)\n    \n    @torch.no_grad()\n    def generate(\n        self,\n        z_neural: torch.Tensor,\n        max_len: int = 200,\n        temperature: float = 1.0,\n        bos_id: int = 1,\n        eos_id: int = 2,\n    ) -> torch.Tensor:\n        \"\"\"\n        Autoregressive generation\n        \n        z_neural: (B, L, d_model)\n        Returns:\n            generated: (B, <= max_len + 1) token IDs (including BOS, until EOS or max_len)\n        \"\"\"\n        B = z_neural.size(0)\n        device = z_neural.device\n        \n        # Start with BOS token (match tokenizer.bos_id)\n        generated = torch.full((B, 1), bos_id, dtype=torch.long, device=device)\n        \n        for _ in range(max_len):\n            # Embed current sequence\n            tgt_emb = self.token_embedding(generated)\n            tgt_emb = self.pos_enc(tgt_emb)\n            \n            # Decode\n            T = generated.size(1)\n            causal_mask = nn.Transformer.generate_square_subsequent_mask(T).to(device)\n            out = self.decoder(tgt_emb, z_neural, tgt_mask=causal_mask)\n            \n            # Get last token logits\n            logits = self.output_proj(out[:, -1, :])  # (B, vocab_size)\n            \n            # Sample next token\n            probs = F.softmax(logits / temperature, dim=-1)\n            next_token = torch.multinomial(probs, num_samples=1)  # (B, 1)\n            \n            # Append\n            generated = torch.cat([generated, next_token], dim=1)\n            \n            # Check for EOS (match tokenizer.eos_id)\n            if (next_token == eos_id).all():\n                break\n        \n        return generated\n\n\n# ============================================================================\n# 3. COMPLETE MODEL - DSD-NLA (SIMPLIFIED)\n# ============================================================================\n\nclass DSDNLA(nn.Module):\n    \"\"\"\n    Simplified DSD-NLA Model\n    \n    Pipeline (current actual behavior):\n    1. Neural features → Neural Encoder → z_neural\n    2. z_neural → Text Decoder → text\n    \n    No diffusion prior, no cross-modal alignment.\n    \"\"\"\n    def __init__(\n        self,\n        n_channels: int = 512,\n        d_model: int = 512,\n        vocab_size: int = 256,\n        n_encoder_layers: int = 8,\n        n_decoder_layers: int = 6,\n        n_heads: int = 8,\n        dropout: float = 0.1,\n    ):\n        super().__init__()\n        \n        # 1. Neural Encoder\n        self.neural_encoder = NeuralEncoder(\n            n_channels=n_channels,\n            d_model=d_model,\n            n_layers=n_encoder_layers,\n            n_heads=n_heads,\n            dropout=dropout\n        )\n        \n        # 2. Text Decoder\n        self.text_decoder = TextDecoder(\n            d_model=d_model,\n            vocab_size=vocab_size,\n            n_layers=n_decoder_layers,\n            n_heads=n_heads,\n            dropout=dropout\n        )\n        \n    def forward(\n        self,\n        neural_features: torch.Tensor,\n        speech_latent=None,          # kept for API compatibility (ignored)\n        target_tokens: torch.Tensor | None = None,\n        training: bool = True,\n    ) -> tuple[torch.Tensor, dict]:\n        \"\"\"\n        neural_features: (B, T, 512)\n        speech_latent: kept for API compatibility (ignored)\n        target_tokens: (B, T) - text tokens for teacher forcing\n        \n        Returns:\n            logits: (B, T, vocab_size)\n            losses: empty dict (no auxiliary losses in this simplified version)\n        \"\"\"\n        # Encode neural\n        z_neural = self.neural_encoder(neural_features)  # (B, L, d_model)\n        \n        # Decode to text\n        logits = self.text_decoder(z_neural, target_tokens=target_tokens)\n        return logits, {}\n    \n    @torch.no_grad()\n    def inference(\n        self,\n        neural_features: torch.Tensor,\n        max_len: int = 200,\n        bos_id: int = 1,\n        eos_id: int = 2,\n    ) -> torch.Tensor:\n        \"\"\"\n        Pure inference: neural → text\n        \n        Returns:\n            tokens: (B, <= max_len + 1) token IDs (including BOS)\n        \"\"\"\n        # Encode\n        z_neural = self.neural_encoder(neural_features)\n        \n        # Decode\n        tokens = self.text_decoder.generate(\n            z_neural,\n            max_len=max_len,\n            bos_id=bos_id,\n            eos_id=eos_id,\n        )\n        return tokens\n\n\n# ============================================================================\n# UTILITIES\n# ============================================================================\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n        \n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() *\n                             (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        pe = pe.unsqueeze(0)\n        self.register_buffer('pe', pe)\n    \n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = x + self.pe[:, :x.size(1)]\n        return self.dropout(x)\n\n\n# ============================================================================\n# EXAMPLE USAGE\n# ============================================================================\n\nif __name__ == \"__main__\":\n    # Config\n    config = {\n        \"n_channels\": 512,\n        \"d_model\": 512,\n        \"vocab_size\": 256,  # Character-level (a-z, A-Z, space, punctuation)\n        \"n_encoder_layers\": 8,\n        \"n_decoder_layers\": 6,\n        \"n_heads\": 8,\n        \"dropout\": 0.1,\n    }\n    \n    model = DSDNLA(**config)\n    \n    # Example forward pass\n    B, T = 4, 200\n    neural_features = torch.randn(B, T, 512)\n    target_tokens = torch.randint(0, 256, (B, 50))\n    \n    # Training-like call\n    logits, losses = model(neural_features, target_tokens=target_tokens, training=True)\n    print(f\"Logits shape: {logits.shape}\")\n    print(f\"Loss dict keys: {losses.keys()}\")\n    \n    # Inference\n    tokens = model.inference(neural_features, max_len=100, bos_id=1, eos_id=2)\n    print(f\"Generated tokens: {tokens.shape}\")\n    \n    print(f\"\\nTotal parameters: {sum(p.numel() for p in model.parameters()):,}\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-10T07:04:38.484498Z","iopub.execute_input":"2025-12-10T07:04:38.485231Z","iopub.status.idle":"2025-12-10T07:04:52.687928Z","shell.execute_reply.started":"2025-12-10T07:04:38.485195Z","shell.execute_reply":"2025-12-10T07:04:52.687206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nDSD-NLA Training Script \nTRUE END-TO-END: Neural → Text (NO phonemes as intermediate!)\n\nCurrent implementation:\n- No diffusion prior\n- No cross-modal alignment\n- Joint training of NeuralEncoder + TextDecoder\n  using character-level text labels (no phonemes).\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport h5py\nimport numpy as np\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport string\nfrom typing import List, Dict, Optional, Any\n\n# ============================================================================#\n# 0. MODEL DEFINITION (must match inference script)                           #\n# ============================================================================#\n\nimport math\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len, dtype=torch.float32).unsqueeze(1)\n        div_term = torch.exp(\n            torch.arange(0, d_model, 2, dtype=torch.float32)\n            * (-math.log(10000.0) / d_model)\n        )\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        pe = pe.unsqueeze(0)  # (1, max_len, d_model)\n        self.register_buffer(\"pe\", pe)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        x: (B, T, d_model)\n        \"\"\"\n        x = x + self.pe[:, : x.size(1)]\n        return self.dropout(x)\n\n\nclass NeuralEncoder(nn.Module):\n    \"\"\"\n    Encode neural signals (512 channels, T timesteps) → latent z_neural\n    \"\"\"\n    def __init__(\n        self,\n        n_channels: int = 512,\n        d_model: int = 512,\n        n_layers: int = 8,\n        n_heads: int = 8,\n        dropout: float = 0.1,\n    ):\n        super().__init__()\n\n        # Input projection - 512 channels → d_model\n        self.input_proj = nn.Sequential(\n            nn.Conv1d(n_channels, d_model, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(d_model),\n            nn.GELU(),\n            nn.Dropout(dropout),\n        )\n\n        # Positional encoding\n        self.pos_enc = PositionalEncoding(d_model, dropout, max_len=5000)\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=n_heads,\n            dim_feedforward=d_model * 4,\n            dropout=dropout,\n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=True,  # Pre-LN for stability\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=n_layers\n        )\n\n        self.layer_norm = nn.LayerNorm(d_model)\n\n    def forward(\n        self,\n        x: torch.Tensor,\n        mask: Optional[torch.Tensor] = None,\n    ) -> torch.Tensor:\n        \"\"\"\n        x: (B, T, 512) neural features (padded with 0.0)\n        mask: (B, T) boolean padding mask BEFORE conv\n              True = PAD, False = valid\n        Returns:\n            z_neural: (B, L, d_model)\n        \"\"\"\n        # Conv projection + downsample 2x\n        x = x.transpose(1, 2)  # (B, 512, T)\n        x = self.input_proj(x)  # (B, d_model, T//2-ish)\n        x = x.transpose(1, 2)  # (B, L, d_model) where L ≈ T//2\n\n        # Downsample mask to match conv stride\n        if mask is not None:\n            # mask: (B, T) → (B, L) via stride 2\n            mask = mask[:, ::2]\n            # In case due to padding we still have length mismatch, clamp\n            if mask.size(1) > x.size(1):\n                mask = mask[:, : x.size(1)]\n            elif mask.size(1) < x.size(1):\n                pad_len = x.size(1) - mask.size(1)\n                pad = torch.ones(\n                    mask.size(0),\n                    pad_len,\n                    dtype=mask.dtype,\n                    device=mask.device,\n                )\n                mask = torch.cat([mask, pad], dim=1)\n\n        # Add positional encoding\n        x = self.pos_enc(x)\n\n        # Transformer encoding (src_key_padding_mask: True = ignore)\n        z_neural = self.transformer(x, src_key_padding_mask=mask)\n        z_neural = self.layer_norm(z_neural)\n\n        return z_neural\n\n\nclass TextDecoder(nn.Module):\n    \"\"\"\n    Decode from neural latent directly to text\n    Not using phoneme intermediate representation!\n    End-to-end mapping from brain → words\n    \"\"\"\n    def __init__(\n        self,\n        d_model: int = 512,\n        vocab_size: int = 256,\n        n_layers: int = 6,\n        n_heads: int = 8,\n        dropout: float = 0.1,\n    ):\n        super().__init__()\n\n        self.d_model = d_model\n        self.vocab_size = vocab_size\n\n        self.token_embedding = nn.Embedding(vocab_size, d_model)\n        self.pos_enc = PositionalEncoding(d_model, dropout)\n\n        decoder_layer = nn.TransformerDecoderLayer(\n            d_model=d_model,\n            nhead=n_heads,\n            dim_feedforward=d_model * 4,\n            dropout=dropout,\n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=True,\n        )\n        self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=n_layers)\n\n        self.output_proj = nn.Linear(d_model, vocab_size)\n\n    def forward(\n        self,\n        z_neural: torch.Tensor,\n        target_tokens: Optional[torch.Tensor] = None,\n        max_len: int = 200,\n    ) -> torch.Tensor:\n        \"\"\"\n        z_neural: (B, L, d_model)\n        target_tokens: (B, T)\n        Returns:\n            (B, T, vocab_size) logits if target_tokens provided\n        \"\"\"\n        if target_tokens is not None:\n            # Teacher forcing\n            tgt_emb = self.token_embedding(target_tokens)\n            tgt_emb = self.pos_enc(tgt_emb)\n\n            T = target_tokens.size(1)\n            causal_mask = nn.Transformer.generate_square_subsequent_mask(T).to(\n                z_neural.device\n            )\n\n            out = self.decoder(tgt_emb, z_neural, tgt_mask=causal_mask)\n            logits = self.output_proj(out)\n            return logits\n\n        # For training we always pass target_tokens, inference is handled by DSDNLA.inference\n        raise RuntimeError(\"TextDecoder.forward called without target_tokens in training mode\")\n\n\nclass DSDNLA(nn.Module):\n    \"\"\"\n    Simplified DSD-NLA Model\n    \n    Pipeline:\n    1. Neural features → Neural Encoder → z_neural\n    2. z_neural → Text Decoder → text\n    No diffusion prior, no cross-modal alignment.\n    \"\"\"\n    def __init__(\n        self,\n        n_channels: int = 512,\n        d_model: int = 512,\n        vocab_size: int = 256,\n        n_encoder_layers: int = 8,\n        n_decoder_layers: int = 6,\n        n_heads: int = 8,\n        dropout: float = 0.1,\n    ):\n        super().__init__()\n\n        self.neural_encoder = NeuralEncoder(\n            n_channels=n_channels,\n            d_model=d_model,\n            n_layers=n_encoder_layers,\n            n_heads=n_heads,\n            dropout=dropout,\n        )\n\n        self.text_decoder = TextDecoder(\n            d_model=d_model,\n            vocab_size=vocab_size,\n            n_layers=n_decoder_layers,\n            n_heads=n_heads,\n            dropout=dropout,\n        )\n\n    def forward(\n        self,\n        neural_features: torch.Tensor,\n        speech_latent=None,  # kept for API compatibility (ignored)\n        target_tokens: Optional[torch.Tensor] = None,\n        neural_mask: Optional[torch.Tensor] = None,\n        training: bool = True,\n    ) -> tuple[torch.Tensor, dict]:\n        \"\"\"\n        neural_features: (B, T, 512)\n        neural_mask: (B, T) bool, True = PAD, False = valid\n        target_tokens: (B, T)\n        Returns:\n            logits: (B, T, vocab_size)\n            losses: empty dict\n        \"\"\"\n        z_neural = self.neural_encoder(neural_features, mask=neural_mask)\n        logits = self.text_decoder(z_neural, target_tokens=target_tokens)\n        return logits, {}\n\n    @torch.no_grad()\n    def inference(\n        self,\n        neural_features: torch.Tensor,\n        max_len: int = 200,\n        bos_id: int = 1,\n        eos_id: int = 2,\n    ) -> torch.Tensor:\n        \"\"\"\n        Pure inference: neural → text\n        Returns:\n            tokens: (B, <= max_len + 1) token IDs (including BOS)\n        \"\"\"\n        # Encode (no padding mask here; inference uses unpadded features)\n        z_neural = self.neural_encoder(neural_features)\n\n        # Greedy generation using the TextDecoder.generate from your inference code.\n        # For training script we don't need full implementation here,\n        # but keep the API intact in case you call it for debugging.\n        from torch.nn import functional as F\n\n        B = z_neural.size(0)\n        device = z_neural.device\n        generated = torch.full(\n            (B, 1), bos_id, dtype=torch.long, device=device\n        )\n\n        for _ in range(max_len):\n            tgt_emb = self.text_decoder.token_embedding(generated)\n            tgt_emb = self.text_decoder.pos_enc(tgt_emb)\n            T = generated.size(1)\n            causal_mask = nn.Transformer.generate_square_subsequent_mask(T).to(device)\n            out = self.text_decoder.decoder(tgt_emb, z_neural, tgt_mask=causal_mask)\n            logits = self.text_decoder.output_proj(out[:, -1, :])\n            probs = F.softmax(logits, dim=-1)\n            next_token = torch.argmax(probs, dim=-1, keepdim=True)\n            generated = torch.cat([generated, next_token], dim=1)\n            if (next_token == eos_id).all():\n                break\n\n        return generated\n\n\n# ============================================================================#\n# 1. TOKENIZER - Character-level                                             #\n# ============================================================================#\n\nclass CharTokenizer:\n    \"\"\"\n    Simple character-level tokenizer\n    0: PAD\n    1: BOS (begin of sequence)\n    2: EOS (end of sequence)\n    3-28: a-z\n    29: space\n    30-?: punctuation\n    \"\"\"\n    def __init__(self):\n        self.pad_id = 0\n        self.bos_id = 1\n        self.eos_id = 2\n\n        # Build vocab\n        self.chars: List[str] = [\"<PAD>\", \"<BOS>\", \"<EOS>\"]\n        self.chars += list(string.ascii_lowercase)  # a-z\n        self.chars += [\" \"]  # space\n        self.chars += list(\"'.,!?-\")  # basic punctuation\n\n        self.char2id: Dict[str, int] = {c: i for i, c in enumerate(self.chars)}\n        self.id2char: Dict[int, str] = {i: c for i, c in enumerate(self.chars)}\n        self.vocab_size: int = len(self.chars)\n\n    def encode(self, text: str) -> torch.Tensor:\n        \"\"\"Text → token IDs (Tensor of shape [T])\"\"\"\n        text = text.lower()\n        ids = [self.bos_id]\n        for c in text:\n            if c in self.char2id:\n                ids.append(self.char2id[c])\n            else:\n                ids.append(self.char2id[\" \"])\n        ids.append(self.eos_id)\n        return torch.tensor(ids, dtype=torch.long)\n\n    def decode(self, ids: torch.Tensor | List[int]) -> str:\n        \"\"\"Token IDs → text (stop at EOS)\"\"\"\n        if isinstance(ids, torch.Tensor):\n            ids_iter = ids.cpu().tolist()\n        else:\n            ids_iter = ids\n\n        chars: List[str] = []\n        for i in ids_iter:\n            if i == self.eos_id:\n                break\n            if i > 2:\n                chars.append(self.id2char.get(i, \" \"))\n        return \"\".join(chars)\n\n\n# ============================================================================#\n# 2. DATASET for Brain-to-Text                                               #\n# ============================================================================#\n\nclass BrainToTextDataset(Dataset):\n    \"\"\"\n    Dataset for DSD-NLA\n    Load neural features + text labels (NO phonemes!)\n    \"\"\"\n    def __init__(\n        self,\n        hdf5_paths: List[str],\n        tokenizer: CharTokenizer,\n        mode: str = \"train\",\n        augment: bool = True,\n    ):\n        self.hdf5_paths = hdf5_paths\n        self.tokenizer = tokenizer\n        self.mode = mode\n        self.augment = augment and (mode == \"train\")\n\n        self.trial_keys: List[tuple[int, str]] = []\n        self.open_files: Dict[int, h5py.File] = {}\n\n        print(f\"Loading {mode} data...\")\n        for i, h5_path in enumerate(tqdm(hdf5_paths)):\n            f = h5py.File(h5_path, \"r\")\n            self.open_files[i] = f\n            for key in f.keys():\n                self.trial_keys.append((i, key))\n\n        print(f\"Loaded {len(self.trial_keys)} trials\")\n\n    def __len__(self) -> int:\n        return len(self.trial_keys)\n\n    def __getitem__(self, idx: int) -> Dict[str, Any]:\n        file_idx, key = self.trial_keys[idx]\n        trial = self.open_files[file_idx][key]\n\n        # Neural features: (T, 512)\n        neural = torch.tensor(trial[\"input_features\"][:], dtype=torch.float32)\n\n        # Text label\n        if \"sentence_label\" in trial.attrs:\n            text = trial.attrs[\"sentence_label\"]\n            tokens = self.tokenizer.encode(text)\n        else:\n            text = None\n            tokens = None\n\n        if self.augment:\n            neural = self.apply_augment(neural)\n\n        return {\n            \"neural\": neural,\n            \"tokens\": tokens,\n            \"text\": text,\n        }\n\n    def apply_augment(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Data augmentation on neural features\"\"\"\n        # Temporal masking\n        if torch.rand(1) < 0.3:\n            mask_len = torch.randint(5, 15, (1,)).item()\n            start = torch.randint(0, max(1, x.size(0) - mask_len), (1,)).item()\n            x[start : start + mask_len] = 0\n\n        # Electrode dropout\n        if torch.rand(1) < 0.2:\n            drop_mask = torch.bernoulli(torch.ones(x.size(1)) * 0.9)\n            x = x * drop_mask\n\n        # Gaussian noise\n        if torch.rand(1) < 0.25:\n            x = x + torch.randn_like(x) * 0.05\n\n        return x\n\n\ndef collate_fn(batch: List[Dict[str, Any]]) -> Dict[str, Any]:\n    \"\"\"Collate with padding for variable-length sequences.\"\"\"\n    neurals = [item[\"neural\"] for item in batch]\n    tokens_list = [item[\"tokens\"] for item in batch if item[\"tokens\"] is not None]\n    texts = [item[\"text\"] for item in batch if item[\"text\"] is not None]\n\n    lengths = [n.size(0) for n in neurals]\n    B = len(neurals)\n    T_max = max(lengths)\n\n    # Pad neural: (B, T_max, 512)\n    neural_padded = nn.utils.rnn.pad_sequence(\n        neurals,\n        batch_first=True,\n        padding_value=0.0,\n    )\n\n    # Build padding mask: True = PAD, False = valid\n    neural_mask = torch.ones(B, T_max, dtype=torch.bool)\n    for i, L in enumerate(lengths):\n        neural_mask[i, :L] = False\n\n    # Pad tokens\n    if tokens_list:\n        tokens_padded = nn.utils.rnn.pad_sequence(\n            tokens_list,\n            batch_first=True,\n            padding_value=0,\n        )\n    else:\n        tokens_padded = None\n\n    return {\n        \"neural\": neural_padded,\n        \"neural_mask\": neural_mask,\n        \"tokens\": tokens_padded,\n        \"texts\": texts,\n    }\n\n\n# ============================================================================#\n# 3. LOSS FUNCTION - Cross-Entropy for text generation (simplified)          #\n# ============================================================================#\n\nclass DSDNLALoss(nn.Module):\n    \"\"\"\n    Simplified loss:\n    - Only text generation loss (cross-entropy)\n    \"\"\"\n    def __init__(self, ignore_index: int = 0):\n        super().__init__()\n        self.ce_loss = nn.CrossEntropyLoss(ignore_index=ignore_index)\n\n    def forward(\n        self,\n        logits: torch.Tensor,\n        target_tokens: Optional[torch.Tensor],\n        model_losses: Optional[Dict[str, torch.Tensor]] = None,\n    ) -> Dict[str, torch.Tensor]:\n        \"\"\"\n        logits: (B, T, vocab_size)\n        target_tokens: (B, T)\n        \"\"\"\n        losses: Dict[str, torch.Tensor] = {}\n\n        if target_tokens is not None:\n            B, T, V = logits.shape\n            logits_flat = logits[:, :-1].reshape(-1, V)\n            target_flat = target_tokens[:, 1:].reshape(-1)\n            text_loss = self.ce_loss(logits_flat, target_flat)\n            losses[\"text\"] = text_loss\n            losses[\"total\"] = text_loss\n        else:\n            zero = torch.tensor(0.0, device=logits.device)\n            losses[\"text\"] = zero\n            losses[\"total\"] = zero\n\n        return losses\n\n\n# ============================================================================#\n# 4. TRAINER                                                                 #\n# ============================================================================#\n\nclass Config:\n    # Data\n    DATA_DIR = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\n\n    # Model\n    n_channels = 512\n    d_model = 512\n    n_encoder_layers = 8\n    n_decoder_layers = 6\n    n_heads = 8\n    dropout = 0.15\n\n    # Training\n    batch_size = 16\n    num_epochs = 50\n    learning_rate = 5e-4\n    weight_decay = 0.01\n    gradient_clip = 1.0\n\n    patience = 5\n    checkpoint_path = \"/kaggle/working/best_dsdnla_model.pt\"\n\n    # Device\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    use_amp = True\n\n    # Logging\n    log_every = 100\n\n\nclass Trainer:\n    def __init__(self, model: nn.Module, tokenizer: CharTokenizer, config: Config):\n        self.model = model\n        self.tokenizer = tokenizer\n        self.config = config\n        self.device = config.device\n\n        self.model = self.model.to(self.device)\n\n        self.optimizer = optim.AdamW(\n            model.parameters(),\n            lr=config.learning_rate,\n            weight_decay=config.weight_decay,\n            betas=(0.9, 0.98),\n            eps=1e-8,\n        )\n\n        self.scheduler: Optional[optim.lr_scheduler._LRScheduler] = None\n\n        self.criterion = DSDNLALoss(ignore_index=0)\n\n        # New AMP API\n        if config.use_amp and self.device == \"cuda\":\n            self.scaler = torch.amp.GradScaler(\"cuda\")\n        else:\n            self.scaler = None\n\n        self.global_step = 0\n        self.best_val_loss = float(\"inf\")\n        self.patience_counter = 0\n\n    def _build_scheduler(self, steps_per_epoch: int):\n        \"\"\"Initialize OneCycleLR once DataLoader length is known.\"\"\"\n        total_steps = max(1, steps_per_epoch * self.config.num_epochs)\n        self.scheduler = optim.lr_scheduler.OneCycleLR(\n            self.optimizer,\n            max_lr=self.config.learning_rate,\n            total_steps=total_steps,\n            pct_start=0.05,\n            anneal_strategy=\"cos\",\n        )\n\n    def train_epoch(self, train_loader: DataLoader, epoch: int) -> float:\n        self.model.train()\n        epoch_losses: List[float] = []\n\n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch}\")\n        for batch in pbar:\n            neural = batch[\"neural\"].to(self.device)\n            neural_mask = batch[\"neural_mask\"].to(self.device)\n            tokens = (\n                batch[\"tokens\"].to(self.device)\n                if batch[\"tokens\"] is not None\n                else None\n            )\n\n            use_autocast = self.scaler is not None\n            autocast_ctx = torch.amp.autocast(\"cuda\", enabled=use_autocast)\n\n            with autocast_ctx:\n                logits, model_losses = self.model(\n                    neural,\n                    speech_latent=None,\n                    target_tokens=tokens,\n                    neural_mask=neural_mask,\n                    training=True,\n                )\n                losses = self.criterion(logits, tokens, model_losses)\n                total_loss = losses[\"total\"]\n\n            self.optimizer.zero_grad(set_to_none=True)\n\n            if self.scaler is not None:\n                self.scaler.scale(total_loss).backward()\n                self.scaler.unscale_(self.optimizer)\n                torch.nn.utils.clip_grad_norm_(\n                    self.model.parameters(), self.config.gradient_clip\n                )\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n            else:\n                total_loss.backward()\n                torch.nn.utils.clip_grad_norm_(\n                    self.model.parameters(), self.config.gradient_clip\n                )\n                self.optimizer.step()\n\n            # Scheduler: ensure optimizer.step() happens BEFORE scheduler.step()\n            if self.scheduler is not None and self.global_step > 0:\n                self.scheduler.step()\n\n            epoch_losses.append(total_loss.item())\n            pbar.set_postfix({\"loss\": f\"{total_loss.item():.4f}\"})\n\n            self.global_step += 1\n\n        return float(np.mean(epoch_losses))\n\n    @torch.no_grad()\n    def validate(self, val_loader: DataLoader) -> float:\n        self.model.eval()\n        val_losses: List[float] = []\n\n        for batch in tqdm(val_loader, desc=\"Validation\"):\n            neural = batch[\"neural\"].to(self.device)\n            neural_mask = batch[\"neural_mask\"].to(self.device)\n            tokens = (\n                batch[\"tokens\"].to(self.device)\n                if batch[\"tokens\"] is not None\n                else None\n            )\n\n            logits, model_losses = self.model(\n                neural,\n                None,\n                tokens,\n                neural_mask=neural_mask,\n                training=False,\n            )\n            losses = self.criterion(logits, tokens, model_losses)\n            val_losses.append(losses[\"total\"].item())\n\n        return float(np.mean(val_losses))\n\n    def save_checkpoint(self, epoch: int):\n        checkpoint = {\n            \"epoch\": epoch,\n            \"global_step\": self.global_step,\n            \"model_state_dict\": self.model.state_dict(),\n            \"optimizer_state_dict\": self.optimizer.state_dict(),\n            \"scheduler_state_dict\": self.scheduler.state_dict()\n            if self.scheduler is not None\n            else None,\n            \"best_val_loss\": self.best_val_loss,\n            \"config\": vars(self.config),\n        }\n\n        save_path = Path(self.config.checkpoint_path)\n        save_path.parent.mkdir(exist_ok=True, parents=True)\n        torch.save(checkpoint, save_path)\n        print(f\"Saved new best model to: {save_path}\")\n\n    def train(self, train_loader: DataLoader, val_loader: DataLoader):\n        print(\"=\" * 50)\n        print(\"Starting DSD-NLA Training\")\n        print(\n            f\"Device: {self.config.device}, Early Stopping Patience: {self.config.patience}\"\n        )\n        print(\"=\" * 50)\n\n        self._build_scheduler(len(train_loader))\n\n        for epoch in range(self.config.num_epochs):\n            train_loss = self.train_epoch(train_loader, epoch)\n            print(f\"\\nEpoch {epoch}: Train Loss = {train_loss:.4f}\")\n\n            val_loss = self.validate(val_loader)\n            print(f\"Epoch {epoch}: Val Loss = {val_loss:.4f}\")\n\n            if val_loss < self.best_val_loss:\n                self.best_val_loss = val_loss\n                self.save_checkpoint(epoch)\n                self.patience_counter = 0\n            else:\n                self.patience_counter += 1\n                print(\n                    f\"No improvement in validation loss. Patience: {self.patience_counter}/{self.config.patience}\"\n                )\n\n            if self.patience_counter >= self.config.patience:\n                print(\n                    f\"\\nEarly stopping triggered after {self.config.patience} epochs with no improvement.\"\n                )\n                print(\n                    f\"Best model saved at {self.config.checkpoint_path} with validation loss {self.best_val_loss:.4f}\"\n                )\n                break\n\n        print(\"\\n\" + \"=\" * 50)\n        print(\"Training Complete!\")\n        print(\"=\" * 50)\n\n\n# ============================================================================#\n# 5. MAIN                                                                    #\n# ============================================================================#\n\ndef main():\n    \"\"\"Main training script\"\"\"\n    from glob import glob\n\n    config = Config()\n\n    tokenizer = CharTokenizer()\n    print(f\"Tokenizer vocab size: {tokenizer.vocab_size}\")\n\n    model = DSDNLA(\n        n_channels=config.n_channels,\n        d_model=config.d_model,\n        vocab_size=tokenizer.vocab_size,\n        n_encoder_layers=config.n_encoder_layers,\n        n_decoder_layers=config.n_decoder_layers,\n        n_heads=config.n_heads,\n        dropout=config.dropout,\n    )\n\n    print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n\n    train_files = sorted(glob(f\"{config.DATA_DIR}/t15.*/data_train.hdf5\"))\n    val_files = sorted(glob(f\"{config.DATA_DIR}/t15.*/data_val.hdf5\"))\n\n    print(f\"\\nFound {len(train_files)} train files\")\n    print(f\"Found {len(val_files)} val files\")\n\n    train_dataset = BrainToTextDataset(\n        train_files, tokenizer, mode=\"train\", augment=True\n    )\n    val_dataset = BrainToTextDataset(\n        val_files, tokenizer, mode=\"val\", augment=False\n    )\n\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        collate_fn=collate_fn,\n        num_workers=2,\n        pin_memory=True,\n    )\n\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config.batch_size,\n        shuffle=False,\n        collate_fn=collate_fn,\n        num_workers=2,\n        pin_memory=True,\n    )\n\n    trainer = Trainer(model, tokenizer, config)\n    trainer.train(train_loader, val_loader)\n\n\nRUN_TRAINING = True\n\nif __name__ == \"__main__\" and RUN_TRAINING:\n    main()\n\nimport os\nprint(\"RUN_TRAINING =\", RUN_TRAINING)\nprint(\"Files in /kaggle/working:\", os.listdir(\"/kaggle/working\"))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T07:04:52.689441Z","iopub.execute_input":"2025-12-10T07:04:52.68982Z","iopub.status.idle":"2025-12-10T07:05:16.516282Z","shell.execute_reply.started":"2025-12-10T07:04:52.689801Z","shell.execute_reply":"2025-12-10T07:05:16.515126Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nDSD-NLA Inference - Generate submission.csv\nTRUE END-TO-END: Neural → Text (NO phonemes!)\n\nInference pipeline:\n- Load trained NeuralEncoder + TextDecoder (DSDNLA)\n- Optionally use beam search + GPT-2 rescoring (if available)\n- Light text cleanup\n- Normalize to strict WER format for submission\n\"\"\"\n\nimport glob\nimport re\nimport string\nfrom pathlib import Path\nfrom typing import List, Dict, Optional\n\nimport h5py\nimport pandas as pd\nimport torch\nfrom tqdm import tqdm\nfrom transformers import GPT2LMHeadModel, GPT2TokenizerFast\n\n\n# ============================================================================\n# TEXT NORMALIZER — Strict WER Format (No punctuation)\n# ============================================================================\n\ndef normalize_for_eval(text: str) -> str:\n    \"\"\"\n    Remove punctuation (except apostrophe),\n    convert to lowercase, normalize whitespace.\n    \"\"\"\n    text = text.lower()\n    text = text.replace(\"’\", \"'\")\n\n    # Remove punctuation except apostrophe\n    text = re.sub(r\"[^a-z0-9'\\s]\", \" \", text)\n\n    # Collapse whitespace\n    text = re.sub(r\"\\s+\", \" \", text).strip()\n\n    return text\n\n\n# ============================================================================\n# CLEANUP FILTER — reduce WER\n# ============================================================================\n\ndef clean_generated_text(text: str) -> str:\n    \"\"\"\n    Heuristics to improve WER without touching the model:\n    - gently fix repeated vowels\n    - fix extra spaces\n    - normalize apostrophe spacing\n    \"\"\"\n    text = text.lower()\n\n    # Safer repetition cleanup: only very long vowel runs, keep two vowels\n    # e.g. \"gooooood\" -> \"good\", but keep consonants intact\n    text = re.sub(r\"([aeiou])\\1{2,}\", r\"\\1\\1\", text)\n\n    # Fix spacing\n    text = \" \".join(text.split())\n\n    # Fix apostrophe spacing\n    text = re.sub(r\"\\s+'\\s*\", \"'\", text)\n    text = re.sub(r\"\\s+'\", \"'\", text)\n\n    return text.strip()\n\n\n# ============================================================================\n# POST-PROCESSING PIPELINE\n# ============================================================================\n\nprint(\"Loading post-processing tools (GPT-2)...\")\n\nDEVICE_LM = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Flag for GPT-2 LM rescoring\nUSE_LM = True\ntry:\n    gpt2_tokenizer = GPT2TokenizerFast.from_pretrained(\"gpt2\")\n    gpt2_model = GPT2LMHeadModel.from_pretrained(\"gpt2\").to(DEVICE_LM)\n    gpt2_model.eval()\n    print(f\"GPT-2 loaded successfully on device: {DEVICE_LM}\")\nexcept Exception as e:\n    print(f\"Warning: Could not load GPT-2, LM rescoring disabled: {e}\")\n    gpt2_tokenizer = None\n    gpt2_model = None\n    USE_LM = False\n\n\n# SpellChecker is DISABLED (identity function) to avoid harming WER.\ndef spell_fix(text: str) -> str:\n    \"\"\"Currently disabled: return text unchanged.\"\"\"\n    return text\n\n\ndef compute_gpt2_perplexity(text: str) -> float:\n    \"\"\"Calculates a pseudo-perplexity score for a text using GPT-2's loss.\"\"\"\n    if not USE_LM:\n        # If LM not available, do not influence ranking\n        return float(\"inf\")\n\n    if not text:  # Handle empty strings\n        return float(\"inf\")\n    try:\n        inputs = gpt2_tokenizer(text, return_tensors=\"pt\").to(gpt2_model.device)\n        with torch.no_grad():\n            outputs = gpt2_model(**inputs, labels=inputs[\"input_ids\"])\n            loss = outputs.loss\n        return loss.item()  # Lower loss is better\n    except Exception as e:\n        print(f\"Warning: Could not compute perplexity due to error: {e}\")\n        return float(\"inf\")\n\n\ndef select_best_candidate_by_lm(candidates: List[str]) -> str:\n    \"\"\"\n    Select the best text from a list of candidates.\n    If GPT-2 is available, use LM perplexity; otherwise fallback to top beam.\n    \"\"\"\n    if not candidates:\n        return \"\"\n    if not USE_LM:\n        # Just take the first (best by beam score)\n        return candidates[0].strip()\n\n    best_score = float(\"inf\")\n    best_text = candidates[0]\n\n    for text in candidates:\n        text_clean = text.strip()\n        ppl = compute_gpt2_perplexity(text_clean)\n        if ppl < best_score:\n            best_score = ppl\n            best_text = text_clean\n\n    return best_text\n\n\n# ============================================================================\n# TOKENIZER\n# ============================================================================\n\nclass CharTokenizer:\n    \"\"\"Character-level tokenizer (same as training).\"\"\"\n    def __init__(self):\n        self.pad_id = 0\n        self.bos_id = 1\n        self.eos_id = 2\n\n        self.chars: List[str] = [\"<PAD>\", \"<BOS>\", \"<EOS>\"]\n        self.chars += list(string.ascii_lowercase)\n        self.chars += [\" \"]\n        self.chars += list(\"'.,!?-\")\n\n        self.char2id: Dict[str, int] = {c: i for i, c in enumerate(self.chars)}\n        self.id2char: Dict[int, str] = {i: c for i, c in enumerate(self.chars)}\n        self.vocab_size: int = len(self.chars)\n\n    def decode(self, ids: torch.Tensor | List[int]) -> str:\n        \"\"\"Token IDs → text\"\"\"\n        if isinstance(ids, torch.Tensor):\n            ids_iter = ids.cpu().tolist()\n        else:\n            ids_iter = ids\n\n        chars: List[str] = []\n        for i in ids_iter:\n            if i == self.eos_id:\n                break\n            if i > 2:  # Skip PAD, BOS, EOS\n                chars.append(self.id2char.get(i, \" \"))\n\n        return \"\".join(chars)\n\n\n# ============================================================================\n# INFERENCE ENGINE\n# ============================================================================\n\nclass InferenceEngine:\n    \"\"\"\n    DSD-NLA Inference Engine\n    Load trained model → generate predictions\n    \"\"\"\n    def __init__(\n        self,\n        model_path: str,\n        device: str = \"cuda\",\n        max_len: int = 200,\n        temperature: float = 0.8,\n    ):\n        self.device = device\n        self.max_len = max_len\n        self.temperature = temperature\n\n        self.tokenizer = CharTokenizer()\n\n        # Resolve and check checkpoint path\n        ckpt_path = Path(model_path)\n        if not ckpt_path.exists():\n            raise FileNotFoundError(f\"Checkpoint not found at {ckpt_path.resolve()}\")\n\n        # Load model\n        print(f\"Loading model from {ckpt_path}...\")\n        checkpoint = torch.load(ckpt_path, map_location=device)\n        config = checkpoint.get(\"config\", {})\n\n        # Create model (NeuralEncoder + TextDecoder only)\n        # NOTE: DSDNLA must be defined / imported elsewhere (simplified version).\n        self.model = DSDNLA(\n            n_channels=config.get(\"n_channels\", 512),\n            d_model=config.get(\"d_model\", 512),\n            vocab_size=self.tokenizer.vocab_size,\n            n_encoder_layers=config.get(\"n_encoder_layers\", 8),\n            n_decoder_layers=config.get(\"n_decoder_layers\", 6),\n            n_heads=config.get(\"n_heads\", 8),\n            dropout=0.0,  # No dropout in inference\n        )\n\n        # Load weights\n        self.model.load_state_dict(checkpoint[\"model_state_dict\"])\n        self.model = self.model.to(device)\n        self.model.eval()\n\n        print(\"Model loaded successfully\")\n        print(f\"  Parameters: {sum(p.numel() for p in self.model.parameters()):,}\")\n\n    @torch.no_grad()\n    def predict(self, neural_features) -> str:\n        \"\"\"\n        Generate text from neural features\n\n        Args:\n            neural_features: (T, 512) numpy array or tensor\n\n        Returns:\n            text: decoded string\n        \"\"\"\n        # Convert to tensor\n        if not isinstance(neural_features, torch.Tensor):\n            neural_features = torch.tensor(neural_features, dtype=torch.float32)\n\n        # Add batch dimension\n        neural_features = neural_features.unsqueeze(0).to(self.device)\n\n        # Generate tokens\n        tokens = self.model.inference(neural_features, max_len=self.max_len)\n\n        # Decode to text\n        tokens = tokens[0]  # Remove batch dimension\n        text = self.tokenizer.decode(tokens)\n\n        # Clean up text\n        text = text.strip()\n        text = \" \".join(text.split())  # Remove multiple spaces\n\n        return text\n\n    @torch.no_grad()\n    def predict_batch(self, neural_batch: List, batch_size: int = 16) -> List[str]:\n        \"\"\"\n        Batch prediction for faster inference\n\n        Args:\n            neural_batch: list of (T_i, 512) arrays\n            batch_size: number of samples per batch\n\n        Returns:\n            texts: list of decoded strings\n        \"\"\"\n        from torch.nn.utils.rnn import pad_sequence\n\n        predictions: List[str] = []\n\n        for i in range(0, len(neural_batch), batch_size):\n            batch_data = neural_batch[i: i + batch_size]\n\n            # Convert to tensors\n            batch_tensors = [torch.tensor(x, dtype=torch.float32) for x in batch_data]\n\n            # Pad\n            padded = pad_sequence(batch_tensors, batch_first=True, padding_value=0.0)\n            padded = padded.to(self.device)\n\n            # Predict sample by sample (keeps logic simple)\n            for j in range(padded.size(0)):\n                neural_single = padded[j: j + 1]\n                tokens = self.model.inference(neural_single, max_len=self.max_len)\n                text = self.tokenizer.decode(tokens[0])\n                text = text.strip()\n                predictions.append(text)\n\n        return predictions\n\n\n# ============================================================================\n# BEAM SEARCH\n# ============================================================================\n\nclass BeamSearchGenerator:\n    \"\"\"\n    Beam search for better text generation\n    \"\"\"\n    def __init__(\n        self,\n        model,\n        tokenizer: CharTokenizer,\n        beam_width: int = 5,\n        max_len: int = 80,\n        length_penalty: float = 0.7,\n        min_len: int = 8,\n    ):\n        self.model = model\n        self.tokenizer = tokenizer\n        self.beam_width = beam_width\n        self.max_len = max_len\n        self.length_penalty = length_penalty\n        self.min_len = min_len\n\n    @torch.no_grad()\n    def generate(self, z_neural: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Beam search generation\n\n        Args:\n            z_neural: (1, L, d_model) encoded neural features\n\n        Returns:\n            best_sequence: (T,) token IDs\n        \"\"\"\n        device = z_neural.device\n\n        # Initialize beams\n        sequences: List[List[int]] = [[self.tokenizer.bos_id]]  # Start with BOS\n        scores: List[float] = [0.0]  # normalized scores\n\n        for _ in range(self.max_len):\n            all_candidates: List[tuple[List[int], float]] = []\n\n            for seq, score in zip(sequences, scores):\n                # Check if ended\n                if seq[-1] == self.tokenizer.eos_id:\n                    all_candidates.append((seq, score))\n                    continue\n\n                # Convert to tensor\n                seq_tensor = torch.tensor([seq], dtype=torch.long, device=device)\n\n                # Get logits from decoder\n                logits = self.model.text_decoder(z_neural, target_tokens=seq_tensor)\n                next_logits = logits[0, -1, :]  # Last token logits\n\n                # Get top-k\n                log_probs = torch.log_softmax(next_logits, dim=-1)\n                topk_probs, topk_ids = torch.topk(log_probs, self.beam_width)\n\n                for prob, idx in zip(topk_probs, topk_ids):\n                    token_id = idx.item()\n                    new_seq = seq + [token_id]\n\n                    # Prevent EOS too early\n                    if token_id == self.tokenizer.eos_id and len(new_seq) < self.min_len:\n                        continue\n\n                    # Length-normalized score\n                    prev_len = len(seq)\n                    new_len = len(new_seq)\n                    # Undo previous normalization then re-normalize\n                    raw_sum = score * (prev_len ** self.length_penalty) if prev_len > 0 else 0.0\n                    raw_sum += prob.item()\n                    new_score = raw_sum / (new_len ** self.length_penalty)\n\n                    all_candidates.append((new_seq, new_score))\n\n            if not all_candidates:\n                # All beams ended early\n                break\n\n            # Select top beam_width candidates\n            all_candidates = sorted(all_candidates, key=lambda x: x[1], reverse=True)\n            sequences = [seq for seq, _ in all_candidates[: self.beam_width]]\n            scores = [score for _, score in all_candidates[: self.beam_width]]\n\n            # Check if all beams ended\n            if all(seq[-1] == self.tokenizer.eos_id for seq in sequences):\n                break\n\n        # Return best sequence\n        best_seq = sequences[0]\n        return torch.tensor(best_seq)\n\n    @torch.no_grad()\n    def generate_candidates(self, z_neural: torch.Tensor) -> List[str]:\n        \"\"\"\n        Beam search generation that returns ALL final candidates as text.\n\n        Args:\n            z_neural: (1, L, d_model) encoded neural features\n\n        Returns:\n            candidate_texts: list[str] (top beam_width candidates).\n        \"\"\"\n        device = z_neural.device\n\n        # Initialize beams\n        sequences: List[List[int]] = [[self.tokenizer.bos_id]]  # Start with BOS\n        scores: List[float] = [0.0]\n\n        for _ in range(self.max_len):\n            all_candidates: List[tuple[List[int], float]] = []\n            all_beams_ended = True\n\n            for seq, score in zip(sequences, scores):\n                if seq[-1] == self.tokenizer.eos_id:\n                    all_candidates.append((seq, score))\n                    continue\n\n                all_beams_ended = False\n\n                seq_tensor = torch.tensor([seq], dtype=torch.long, device=device)\n                logits = self.model.text_decoder(z_neural, target_tokens=seq_tensor)\n                next_logits = logits[0, -1, :]\n                log_probs = torch.log_softmax(next_logits, dim=-1)\n                topk_probs, topk_ids = torch.topk(log_probs, self.beam_width)\n\n                for prob, idx in zip(topk_probs, topk_ids):\n                    token_id = idx.item()\n                    new_seq = seq + [token_id]\n\n                    # Prevent EOS too early\n                    if token_id == self.tokenizer.eos_id and len(new_seq) < self.min_len:\n                        continue\n\n                    prev_len = len(seq)\n                    new_len = len(new_seq)\n                    raw_sum = score * (prev_len ** self.length_penalty) if prev_len > 0 else 0.0\n                    raw_sum += prob.item()\n                    new_score = raw_sum / (new_len ** self.length_penalty)\n\n                    all_candidates.append((new_seq, new_score))\n\n            if all_beams_ended or not all_candidates:\n                break\n\n            all_candidates = sorted(all_candidates, key=lambda x: x[1], reverse=True)\n            sequences = [seq for seq, _ in all_candidates[: self.beam_width]]\n            scores = [score for _, score in all_candidates[: self.beam_width]]\n\n        candidate_texts = [self.tokenizer.decode(seq) for seq in sequences]\n        return candidate_texts\n\n\n# ============================================================================\n# SUBMISSION GENERATOR\n# ============================================================================\n\ndef generate_submission(\n    model_path: str,\n    test_data_dir: str,\n    output_path: str = \"submission.csv\",\n    device: str = \"cuda\",\n    use_beam_search: bool = False,\n    beam_width: int = 5,\n) -> pd.DataFrame:\n    # Initialize inference engine\n    engine = InferenceEngine(model_path=model_path, device=device)\n\n    # Initialize the beam search generator if requested\n    beam_generator: Optional[BeamSearchGenerator] = None\n    if use_beam_search:\n        beam_generator = BeamSearchGenerator(\n            engine.model,\n            engine.tokenizer,\n            beam_width=beam_width,\n        )\n\n    test_files = sorted(glob.glob(f\"{test_data_dir}/t15.*/data_test.hdf5\"))\n    print(f\"Found {len(test_files)} test files.\")\n\n    all_predictions: List[str] = []\n\n    print(\"\\nStarting prediction generation with advanced post-processing...\")\n    for file_path in tqdm(test_files, desc=\"Files\"):\n        with h5py.File(file_path, \"r\") as f:\n            keys = sorted(f.keys())\n            for key in tqdm(\n                keys,\n                desc=f\"Trials in {Path(file_path).parent.name}\",\n                leave=False,\n            ):\n                neural_features = f[key][\"input_features\"][:]\n\n                if use_beam_search and beam_generator is not None:\n                    # 1. Get ALL candidates from beam search\n                    neural_tensor = torch.tensor(neural_features).unsqueeze(0).to(device)\n                    z_neural = engine.model.neural_encoder(neural_tensor)\n                    beam_candidates = beam_generator.generate_candidates(z_neural)\n\n                    # 2. Rescore candidates with GPT-2 (if available) to select the most fluent one\n                    best_text_from_lm = select_best_candidate_by_lm(beam_candidates)\n\n                    # 3. (Optionally) apply spell correction — currently identity\n                    final_text = spell_fix(best_text_from_lm)\n                else:  # Greedy decoding\n                    raw_text = engine.predict(neural_features)\n                    final_text = spell_fix(raw_text)\n\n                # 4. Normalize the final text for the submission format\n                final_text = clean_generated_text(final_text)\n                normalized_text = normalize_for_eval(final_text)\n                all_predictions.append(normalized_text)\n\n    print(f\"\\nGenerated {len(all_predictions)} predictions.\")\n    submission_df = pd.DataFrame({\"id\": range(len(all_predictions)), \"text\": all_predictions})\n    submission_df.to_csv(output_path, index=False)\n    print(f\"Submission file saved to {output_path}\")\n\n    print(\"\\nFirst 10 predictions:\")\n    print(submission_df.head(10))\n\n    return submission_df\n\n\n# ============================================================================\n# FINAL KAGGLE-SAFE EXECUTION\n# ============================================================================\n\nmodel_path = \"/kaggle/working/best_dsdnla_model.pt\"\ntest_data_dir = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\noutput_path = \"/kaggle/working/submission.csv\"\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nuse_beam_search = True\nbeam_width = 10\n\nsubmission_df = generate_submission(\n    model_path=model_path,\n    test_data_dir=test_data_dir,\n    output_path=output_path,\n    device=device,\n    use_beam_search=use_beam_search,\n    beam_width=beam_width,\n)\n\nprint(\"\\n✅ Kaggle will now detect this file:\")\nprint(output_path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T07:05:16.518552Z","iopub.status.idle":"2025-12-10T07:05:16.518871Z","shell.execute_reply.started":"2025-12-10T07:05:16.518707Z","shell.execute_reply":"2025-12-10T07:05:16.518721Z"}},"outputs":[],"execution_count":null}]}