{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.13"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"},{"sourceId":12491966,"sourceType":"datasetVersion","datasetId":7875034},{"sourceId":13773948,"sourceType":"datasetVersion","datasetId":8766624,"isSourceIdPinned":true},{"sourceId":668897,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":506487,"modelId":521271}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":113.715298,"end_time":"2025-10-31T19:04:06.789773","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-10-31T19:02:13.074475","version":"2.6.0"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"199a850c0f2a4355a1d9b0d86af05ec8":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"26ee45f8ecd44d208a0598738764c4e4":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_40a3f84b91a143e0a35573f14833c804","IPY_MODEL_9496b1f87be64e61b634ba71ce4e43e9","IPY_MODEL_2a2d5f08e18f4febb1ffa5d0cbe25ffb"],"layout":"IPY_MODEL_dde92d9903cd403f9758b8f2c4a54aab","tabbable":null,"tooltip":null}},"2a2d5f08e18f4febb1ffa5d0cbe25ffb":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HTMLView","description":"","description_allow_html":false,"layout":"IPY_MODEL_a7c8e5ac3a184b648bd406f50ec48d7a","placeholder":"​","style":"IPY_MODEL_e68644462fe84146bfda2e81d3a90b66","tabbable":null,"tooltip":null,"value":" 41/41 [00:00&lt;00:00, 49.00it/s]"}},"40a3f84b91a143e0a35573f14833c804":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HTMLView","description":"","description_allow_html":false,"layout":"IPY_MODEL_da26d8d0f44a4c8faac3df7f7b698efc","placeholder":"​","style":"IPY_MODEL_f59879cdee45426a977a988e1bff1975","tabbable":null,"tooltip":null,"value":"Loading files: 100%"}},"4d0f28e9bff140e29f5e325726fa0a02":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"72241ce289d94db581894b6bd6576845":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_b41d1224cda0498eadd494d46ba6db59","IPY_MODEL_7e6fba635f654a4daa64e7193baf37cf","IPY_MODEL_90d0b937be0e4167810e9719a0f2a372"],"layout":"IPY_MODEL_199a850c0f2a4355a1d9b0d86af05ec8","tabbable":null,"tooltip":null}},"7e6fba635f654a4daa64e7193baf37cf":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"ProgressView","bar_style":"success","description":"","description_allow_html":false,"layout":"IPY_MODEL_4d0f28e9bff140e29f5e325726fa0a02","max":46,"min":0,"orientation":"horizontal","style":"IPY_MODEL_aea1159cc55240e4ad75b0ae70d96014","tabbable":null,"tooltip":null,"value":46}},"90d0b937be0e4167810e9719a0f2a372":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HTMLView","description":"","description_allow_html":false,"layout":"IPY_MODEL_c531ee5da0ae4344aa925d40876ba304","placeholder":"​","style":"IPY_MODEL_d000c26a3d894aa8b52c3be9ae907d21","tabbable":null,"tooltip":null,"value":" 46/46 [00:36&lt;00:00,  1.34it/s]"}},"9496b1f87be64e61b634ba71ce4e43e9":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"ProgressView","bar_style":"success","description":"","description_allow_html":false,"layout":"IPY_MODEL_c9313891222c4ae1b732ecd5c80df4cc","max":41,"min":0,"orientation":"horizontal","style":"IPY_MODEL_b68ffe36e8ac452e942815477c4d5ad4","tabbable":null,"tooltip":null,"value":41}},"a7c8e5ac3a184b648bd406f50ec48d7a":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"a8ffada8cef745a4a6bc5c0960889cfc":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"aea1159cc55240e4ad75b0ae70d96014":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"b41d1224cda0498eadd494d46ba6db59":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HTMLView","description":"","description_allow_html":false,"layout":"IPY_MODEL_a8ffada8cef745a4a6bc5c0960889cfc","placeholder":"​","style":"IPY_MODEL_e223c7d3c5ec46e8ba64df8c8b24294b","tabbable":null,"tooltip":null,"value":"Decoding Test Set: 100%"}},"b68ffe36e8ac452e942815477c4d5ad4":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"c531ee5da0ae4344aa925d40876ba304":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"c9313891222c4ae1b732ecd5c80df4cc":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"d000c26a3d894aa8b52c3be9ae907d21":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","background":null,"description_width":"","font_size":null,"text_color":null}},"da26d8d0f44a4c8faac3df7f7b698efc":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"dde92d9903cd403f9758b8f2c4a54aab":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"e223c7d3c5ec46e8ba64df8c8b24294b":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","background":null,"description_width":"","font_size":null,"text_color":null}},"e68644462fe84146bfda2e81d3a90b66":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","background":null,"description_width":"","font_size":null,"text_color":null}},"f59879cdee45426a977a988e1bff1975":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","background":null,"description_width":"","font_size":null,"text_color":null}}},"version_major":2,"version_minor":0}}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nDSD-NLA: Diffusion Speech Decoder + Neural Latent Alignment\nImplementation -  END-TO-END\n\nArchitechture:\n1. Neural Encoder (Transformer/Conformer) → z_neural\n2. Diffusion Prior learning from ASR data → aligns latent spaces\n3. Cross-modal Alignment (Contrastive + EM-style) → \"brain-to-unit translator\"\n4. Text Decoder (Transformer LM) → text (not through phoneme)\n\nNot using CTC! end-to-end from neural → text\n\"\"\"\n!pip install pyspellchecker\n!pip install language-tool-python\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport math\nfrom typing import Optional, Tuple, List\nimport numpy as np\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=512, d_model=512, n_layers=8, n_heads=8, dropout=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, mask=None):\n        \"\"\"\n        x: (B, T, 512) neural features\n        Returns: (B, L, d_model) neural latent\n        \"\"\"\n        B, T, C = x.shape\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. DIFFUSION PRIOR - Learning from ASR datasets (LibriSpeech, Switchboard)\n# ============================================================================\n\nclass DiffusionPrior(nn.Module):\n    \"\"\"\n    Diffusion model learn speech prior from ASR data\n    \n    Training: Pretrain on speech units (HuBERT/wav2vec2 units)\n    Inference: Used as a prior to align neural latent with speech distribution\n    \n    KEY INNOVATION: Diffusion is not used to generate audio,\n    it is used to ALIGN latent spaces!\n    \"\"\"\n    def __init__(self, d_model=512, n_steps=1000, beta_start=1e-4, beta_end=0.02):\n        super().__init__()\n        \n        self.d_model = d_model\n        self.n_steps = n_steps\n        \n        # Noise schedule (linear or cosine)\n        betas = torch.linspace(beta_start, beta_end, n_steps)\n        alphas = 1.0 - betas\n        alphas_cumprod = torch.cumprod(alphas, dim=0)\n        \n        self.register_buffer('betas', betas)\n        self.register_buffer('alphas', alphas)\n        self.register_buffer('alphas_cumprod', alphas_cumprod)\n        self.register_buffer('sqrt_alphas_cumprod', torch.sqrt(alphas_cumprod))\n        self.register_buffer('sqrt_one_minus_alphas_cumprod', torch.sqrt(1.0 - alphas_cumprod))\n        \n        # Denoising network (Transformer-based U-Net style)\n        self.time_embed = nn.Sequential(\n            SinusoidalPosEmb(d_model),\n            nn.Linear(d_model, d_model * 4),\n            nn.GELU(),\n            nn.Linear(d_model * 4, d_model)\n        )\n        \n        # Denoising Transformer\n        self.denoiser = nn.ModuleList([\n            nn.TransformerEncoderLayer(\n                d_model=d_model,\n                nhead=8,\n                dim_feedforward=d_model * 4,\n                dropout=0.1,\n                activation='gelu',\n                batch_first=True,\n                norm_first=True\n            ) for _ in range(6)\n        ])\n        \n        self.out_proj = nn.Linear(d_model, d_model)\n        \n    def q_sample(self, x_0, t, noise=None):\n        \"\"\"\n        Forward diffusion: add noise into x_0\n        x_0: (B, L, d_model) clean latent\n        t: (B,) timestep\n        \"\"\"\n        if noise is None:\n            noise = torch.randn_like(x_0)\n        \n        sqrt_alpha_t = self.sqrt_alphas_cumprod[t].view(-1, 1, 1)\n        sqrt_one_minus_alpha_t = self.sqrt_one_minus_alphas_cumprod[t].view(-1, 1, 1)\n        \n        x_t = sqrt_alpha_t * x_0 + sqrt_one_minus_alpha_t * noise\n        return x_t, noise\n    \n    def forward(self, x_t, t):\n        \"\"\"\n        Reverse diffusion: predict noise from x_t\n        x_t: (B, L, d_model) noisy latent\n        t: (B,) timestep\n        Returns: predicted noise\n        \"\"\"\n        B, L, D = x_t.shape\n        \n        # Time embedding\n        t_emb = self.time_embed(t).unsqueeze(1)  # (B, 1, d_model)\n        \n        # Add time info to input\n        x = x_t + t_emb\n        \n        # Denoise\n        for layer in self.denoiser:\n            x = layer(x)\n        \n        noise_pred = self.out_proj(x)\n        return noise_pred\n    \n    @torch.no_grad()\n    def sample(self, z_neural, n_steps=None):\n        \"\"\"\n        Sample from prior, conditioned on neural latent\n        Using in inference to \"clean up\" neural latent\n        \"\"\"\n        if n_steps is None:\n            n_steps = self.n_steps\n        \n        B, L, D = z_neural.shape\n        \n        # Start from noise\n        x = torch.randn_like(z_neural)\n        \n        # Reverse diffusion\n        for t in reversed(range(n_steps)):\n            t_batch = torch.full((B,), t, device=x.device, dtype=torch.long)\n            \n            # Predict noise\n            noise_pred = self.forward(x, t_batch)\n            \n            # Remove noise\n            alpha_t = self.alphas[t]\n            alpha_cumprod_t = self.alphas_cumprod[t]\n            beta_t = self.betas[t]\n            \n            # Mean\n            x = (x - beta_t / torch.sqrt(1 - alpha_cumprod_t) * noise_pred) / torch.sqrt(alpha_t)\n            \n            # Add noise (except last step)\n            if t > 0:\n                noise = torch.randn_like(x)\n                x = x + torch.sqrt(beta_t) * noise\n        \n        return x\n\n\nclass SinusoidalPosEmb(nn.Module):\n    \"\"\"Sinusoidal timestep embedding\"\"\"\n    def __init__(self, dim):\n        super().__init__()\n        self.dim = dim\n    \n    def forward(self, t):\n        device = t.device\n        half_dim = self.dim // 2\n        emb = math.log(10000) / (half_dim - 1)\n        emb = torch.exp(torch.arange(half_dim, device=device) * -emb)\n        emb = t[:, None] * emb[None, :]\n        emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)\n        return emb\n\n\n# ============================================================================\n# 3. CROSS-MODAL ALIGNMENT - \"Brain-to-Unit Translator\"\n# ============================================================================\n\nclass CrossModalAligner(nn.Module):\n    \"\"\"\n    Align neural latent z_neural with speech latent z_speech\n    \n    Using:\n    - Contrastive learning (InfoNCE)\n    - EM-style alternation to learn alignment\n    - Mutual Information maximization\n    \"\"\"\n    def __init__(self, d_model=512, temperature=0.07):\n        super().__init__()\n        \n        self.temperature = temperature\n        \n        # Projection heads\n        self.proj_neural = nn.Sequential(\n            nn.Linear(d_model, d_model),\n            nn.LayerNorm(d_model),\n            nn.GELU(),\n            nn.Linear(d_model, d_model)\n        )\n        \n        self.proj_speech = nn.Sequential(\n            nn.Linear(d_model, d_model),\n            nn.LayerNorm(d_model),\n            nn.GELU(),\n            nn.Linear(d_model, d_model)\n        )\n        \n        # Alignment predictor (soft attention)\n        self.alignment_net = nn.MultiheadAttention(\n            embed_dim=d_model,\n            num_heads=8,\n            dropout=0.1,\n            batch_first=True\n        )\n        \n    def info_nce_loss(self, z_neural, z_speech):\n        \"\"\"\n        Contrastive loss between neural and speech latent\n        \"\"\"\n        B, L_n, D = z_neural.shape\n        _, L_s, _ = z_speech.shape\n        \n        # Project\n        z_n = self.proj_neural(z_neural)  # (B, L_n, D)\n        z_s = self.proj_speech(z_speech)  # (B, L_s, D)\n        \n        # L2 normalize\n        z_n = F.normalize(z_n, dim=-1)\n        z_s = F.normalize(z_s, dim=-1)\n        \n        # Compute similarity (with alignment)\n        # Using attention to find correspondence\n        aligned_z_s, attn_weights = self.alignment_net(z_n, z_s, z_s)\n        \n        # Contrastive loss\n        # Positive pairs: aligned representations\n        pos_sim = (z_n * aligned_z_s).sum(dim=-1)  # (B, L_n)\n        \n        # Negative pairs: all other representations\n        # Flatten\n        z_n_flat = z_n.reshape(B * L_n, D)\n        z_s_flat = z_s.reshape(B * L_s, D)\n        \n        logits = torch.matmul(z_n_flat, z_s_flat.T) / self.temperature  # (B*L_n, B*L_s)\n        \n        # InfoNCE loss\n        labels = torch.arange(B * L_n, device=logits.device) % (B * L_s)\n        loss = F.cross_entropy(logits, labels)\n        \n        return loss, attn_weights\n    \n    def forward(self, z_neural, z_speech):\n        \"\"\"\n        Returns alignment loss and aligned speech representation\n        \"\"\"\n        loss, attn_weights = self.info_nce_loss(z_neural, z_speech)\n        \n        # Get aligned speech\n        aligned_speech, _ = self.alignment_net(z_neural, z_speech, z_speech)\n        \n        return {\n            'loss': loss,\n            'aligned_speech': aligned_speech,\n            'attn_weights': attn_weights\n        }\n\n\n# ============================================================================\n# 4. 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=512, vocab_size=256, n_layers=6, n_heads=8, dropout=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(self, z_neural, target_tokens=None, max_len=200):\n        \"\"\"\n        z_neural: (B, L, d_model) encoded neural features\n        target_tokens: (B, T) ground truth tokens (during training)\n        \n        Returns: (B, T, vocab_size) logits\n        \"\"\"\n        B = z_neural.size(0)\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            \n            return logits\n        else:\n            # Autoregressive generation (inference)\n            return self.generate(z_neural, max_len)\n    \n    @torch.no_grad()\n    def generate(self, z_neural, max_len=200, temperature=1.0):\n        \"\"\"\n        Autoregressive generation\n        \"\"\"\n        B = z_neural.size(0)\n        device = z_neural.device\n        \n        # Start with BOS token (assume vocab[0] = BOS)\n        generated = torch.zeros(B, 1, 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 (assume vocab[1] = EOS)\n            if (next_token == 1).all():\n                break\n        \n        return generated\n\n\n# ============================================================================\n# 5. COMPLETE MODEL - DSD-NLA\n# ============================================================================\n\nclass DSDNLA(nn.Module):\n    \"\"\"\n    Complete DSD-NLA Model\n    \n    Pipeline:\n    1. Neural features → Neural Encoder → z_neural\n    2. (During training) Speech units → Diffusion Prior → z_speech_clean\n    3. Alignment: z_neural ↔ z_speech using CrossModalAligner\n    4. z_neural → Text Decoder → text\n    \n    KEY: Using Diffusion prior only in training to teach alignment!\n    Inference: z_neural → Text Decoder directly\n    \"\"\"\n    def __init__(\n        self,\n        n_channels=512,\n        d_model=512,\n        vocab_size=256,\n        n_encoder_layers=8,\n        n_decoder_layers=6,\n        n_heads=8,\n        dropout=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. Diffusion Prior \n        self.diffusion_prior = DiffusionPrior(d_model=d_model)\n        \n        # 3. Cross-modal Aligner\n        self.aligner = CrossModalAligner(d_model=d_model)\n        \n        # 4. 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(self, neural_features, speech_latent=None, target_tokens=None, training=True):\n        \"\"\"\n        neural_features: (B, T, 512)\n        speech_latent: (B, L, d_model) - from ASR data, only during training\n        target_tokens: (B, T) - text tokens for teacher forcing\n        \n        Returns:\n            logits: (B, T, vocab_size)\n            losses: dict of auxiliary losses\n        \"\"\"\n        # 1. Encode neural\n        z_neural = self.neural_encoder(neural_features)  # (B, L, d_model)\n        \n        losses = {}\n        \n        if training and speech_latent is not None:\n            # 2. Diffusion loss (learn speech prior)\n            B = z_neural.size(0)\n            t = torch.randint(0, self.diffusion_prior.n_steps, (B,), device=z_neural.device)\n            \n            # Add noise to speech latent\n            z_noisy, noise_true = self.diffusion_prior.q_sample(speech_latent, t)\n            \n            # Predict noise\n            noise_pred = self.diffusion_prior(z_noisy, t)\n            \n            diffusion_loss = F.mse_loss(noise_pred, noise_true)\n            losses['diffusion'] = diffusion_loss\n            \n            # 3. Alignment loss\n            alignment_out = self.aligner(z_neural, speech_latent)\n            losses['alignment'] = alignment_out['loss']\n            \n            # Use aligned speech for decoding (optional)\n            # z_for_decode = alignment_out['aligned_speech']\n            z_for_decode = z_neural  # Or use aligned version\n        else:\n            z_for_decode = z_neural\n        \n        # 4. Decode to text\n        logits = self.text_decoder(z_for_decode, target_tokens)\n        \n        return logits, losses\n    \n    @torch.no_grad()\n    def inference(self, neural_features, max_len=200):\n        \"\"\"\n        Pure inference: neural → text\n        NO DIFFUSION, NO ALIGNMENT\n        \"\"\"\n        # Encode\n        z_neural = self.neural_encoder(neural_features)\n        \n        # Optional: clean with diffusion prior \n        # z_clean = self.diffusion_prior.sample(z_neural, n_steps=50)\n        \n        # Decode\n        tokens = self.text_decoder.generate(z_neural, max_len=max_len)\n        \n        return tokens\n\n\n# ============================================================================\n# UTILITIES\n# ============================================================================\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, dropout=0.1, max_len=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() * (-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):\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    speech_latent = torch.randn(B, T//2, 512)  # From ASR\n    target_tokens = torch.randint(0, 256, (B, 50))\n    \n    # Training\n    logits, losses = model(neural_features, speech_latent, target_tokens, training=True)\n    print(f\"Logits shape: {logits.shape}\")\n    print(f\"Losses: {losses.keys()}\")\n    \n    # Inference\n    tokens = model.inference(neural_features, max_len=100)\n    print(f\"Generated tokens: {tokens.shape}\")\n    \n    print(f\"\\nTotal parameters: {sum(p.numel() for p in model.parameters()):,}\")","metadata":{"papermill":{"duration":0.003013,"end_time":"2025-10-31T19:04:04.068871","exception":false,"start_time":"2025-10-31T19:04:04.065858","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-12-01T16:18:21.924146Z","iopub.execute_input":"2025-12-01T16:18:21.924732Z","iopub.status.idle":"2025-12-01T16:18:38.818623Z","shell.execute_reply.started":"2025-12-01T16:18:21.924709Z","shell.execute_reply":"2025-12-01T16:18:38.817671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nDSD-NLA Training Script \nTRUE END-TO-END: Neural → Text (NO phonemes as intermediate!)\n\nTraining Strategy:\nStage 1: Pretrain Diffusion Prior trên ASR data (speech units)\nStage 2: Train Alignment + Decoder với supervision từ text labels\nStage 3: Fine-tune end-to-end\n\nText labels are tokenized to characters/BPE, NOT THROUGH 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\n\n# ============================================================================\n# 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-39: 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 = ['<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 = {c: i for i, c in enumerate(self.chars)}\n        self.id2char = {i: c for i, c in enumerate(self.chars)}\n        self.vocab_size = len(self.chars)\n    \n    def encode(self, text):\n        \"\"\"Text → token IDs\"\"\"\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                # Unknown chars → space\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):\n        \"\"\"Token IDs → text\"\"\"\n        chars = []\n        for i in ids:\n            if i == self.eos_id:\n                break\n            if i > 2:  # Skip PAD, BOS, EOS\n                chars.append(self.id2char[i.item()])\n        return ''.join(chars)\n\n\n# ============================================================================\n# DATASET for Brain-to-Text\n# ============================================================================\n\nclass BrainToTextDataset(Dataset):\n    \"\"\"\n    Dataset cho DSD-NLA\n    Load neural features + text labels (NO phonemes!)\n    \"\"\"\n    def __init__(self, hdf5_paths, tokenizer, mode='train', augment=True):\n        self.hdf5_paths = hdf5_paths\n        self.tokenizer = tokenizer\n        self.mode = mode\n        self.augment = augment and (mode == 'train')\n        \n        # Load trials\n        self.trial_keys = []\n        self.open_files = {}\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):\n        return len(self.trial_keys)\n    \n    def __getitem__(self, idx):\n        file_idx, key = self.trial_keys[idx]\n        trial = self.open_files[file_idx][key]\n        \n        # Neural features\n        neural = torch.tensor(trial['input_features'][:], dtype=torch.float32)\n        \n        # Text label (NOT phonemes!)\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        # Augmentation\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):\n        \"\"\"Data augmentation\"\"\"\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):\n    \"\"\"Collate with padding\"\"\"\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    # Pad neural\n    neural_padded = nn.utils.rnn.pad_sequence(neurals, batch_first=True, padding_value=0.0)\n    \n    # Pad tokens\n    if tokens_list:\n        tokens_padded = nn.utils.rnn.pad_sequence(tokens_list, batch_first=True, padding_value=0)\n    else:\n        tokens_padded = None\n    \n    return {\n        'neural': neural_padded,\n        'tokens': tokens_padded,\n        'texts': texts\n    }\n\n\n# ============================================================================\n# LOSS FUNCTION - Cross-Entropy cho text generation\n# ============================================================================\n\nclass DSDNLALoss(nn.Module):\n    \"\"\"\n    Combined loss:\n    1. Text generation loss (cross-entropy)\n    2. Diffusion loss (MSE)\n    3. Alignment loss (InfoNCE)\n    \"\"\"\n    def __init__(self, text_weight=1.0, diffusion_weight=0.5, alignment_weight=0.3):\n        super().__init__()\n        self.text_weight = text_weight\n        self.diffusion_weight = diffusion_weight\n        self.alignment_weight = alignment_weight\n        self.ce_loss = nn.CrossEntropyLoss(ignore_index=0)  # Ignore PAD\n    \n    def forward(self, logits, target_tokens, model_losses):\n        \"\"\"\n        logits: (B, T, vocab_size)\n        target_tokens: (B, T)\n        model_losses: dict from model forward\n        \"\"\"\n        total_loss = 0\n        losses = {}\n        \n        # 1. Text generation loss\n        if target_tokens is not None:\n            B, T, V = logits.shape\n            logits_flat = logits[:, :-1].reshape(-1, V)  # Shift right\n            target_flat = target_tokens[:, 1:].reshape(-1)  # Predict next\n            \n            text_loss = self.ce_loss(logits_flat, target_flat)\n            losses['text'] = text_loss\n            total_loss += self.text_weight * text_loss\n        \n        # 2. Diffusion loss\n        if 'diffusion' in model_losses:\n            diffusion_loss = model_losses['diffusion']\n            losses['diffusion'] = diffusion_loss\n            total_loss += self.diffusion_weight * diffusion_loss\n        \n        # 3. Alignment loss\n        if 'alignment' in model_losses:\n            alignment_loss = model_losses['alignment']\n            losses['alignment'] = alignment_loss\n            total_loss += self.alignment_weight * alignment_loss\n        \n        losses['total'] = total_loss\n        return losses\n\n\n# ============================================================================\n# 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 = \"best_dsdnla_model.pt\"\n    \n    # Loss weights\n    text_weight = 1.0\n    diffusion_weight = 0.5\n    alignment_weight = 0\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, tokenizer, 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 = optim.lr_scheduler.OneCycleLR(\n            self.optimizer,\n            max_lr=config.learning_rate,\n            total_steps=config.num_epochs * 1000, \n            pct_start=0.05,\n            anneal_strategy='cos'\n        )\n        \n        self.criterion = DSDNLALoss(\n            text_weight=config.text_weight,\n            diffusion_weight=config.diffusion_weight,\n            alignment_weight=config.alignment_weight\n        )\n        \n        self.scaler = torch.cuda.amp.GradScaler() if config.use_amp else None\n        \n        self.global_step = 0\n        self.best_val_loss = float('inf')\n        \n        self.patience_counter = 0\n\n    def train_epoch(self, train_loader, epoch):\n        self.model.train()\n        epoch_losses = []\n        \n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch}\")\n        for batch in pbar:\n            neural = batch['neural'].to(self.device)\n            tokens = batch['tokens'].to(self.device) if batch['tokens'] is not None else None\n            \n            with torch.cuda.amp.autocast(enabled=self.config.use_amp):\n                logits, model_losses = self.model(\n                    neural, \n                    speech_latent=None,\n                    target_tokens=tokens,\n                    training=True\n                )\n                losses = self.criterion(logits, tokens, model_losses)\n                total_loss = losses['total']\n            \n            self.optimizer.zero_grad()\n            \n            if self.scaler:\n                self.scaler.scale(total_loss).backward()\n                self.scaler.unscale_(self.optimizer)\n                torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.config.gradient_clip)\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n            else:\n                total_loss.backward()\n                torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.config.gradient_clip)\n                self.optimizer.step()\n            \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        \n        return np.mean(epoch_losses)\n    \n    @torch.no_grad()\n    def validate(self, val_loader):\n        self.model.eval()\n        val_losses = []\n        \n        for batch in tqdm(val_loader, desc=\"Validation\"):\n            neural = batch['neural'].to(self.device)\n            tokens = batch['tokens'].to(self.device) if batch['tokens'] is not None else None\n            \n            logits, model_losses = self.model(neural, None, tokens, training=False)\n            losses = self.criterion(logits, tokens, model_losses)\n            \n            val_losses.append(losses['total'].item())\n        \n        return np.mean(val_losses)\n    \n    def save_checkpoint(self, epoch):\n        \"\"\"Save the best model checkpoint\"\"\"\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            '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, val_loader):\n        print(\"=\"*50)\n        print(\"Starting DSD-NLA Training\")\n        print(f\"Device: {self.config.device}, Early Stopping Patience: {self.config.patience}\")\n        print(\"=\"*50)\n        \n        for epoch in range(self.config.num_epochs):\n            # Train\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(f\"No improvement in validation loss. Patience: {self.patience_counter}/{self.config.patience}\")\n\n            if self.patience_counter >= self.config.patience:\n                print(f\"\\nEarly stopping triggered after {self.config.patience} epochs with no improvement.\")\n                print(f\"Best model saved at {self.config.checkpoint_path} with validation loss {self.best_val_loss:.4f}\")\n                break \n        \n        print(\"\\n\" + \"=\"*50)\n        print(\"Training Complete!\")\n        print(\"=\"*50)\n\n\n# ============================================================================\n# MAIN\n# ============================================================================\n\ndef main():\n    \"\"\"Main training script\"\"\"\n    from glob import glob\n    \n    config = Config()\n    \n    # Initialize tokenizer\n    tokenizer = CharTokenizer()\n    print(f\"Tokenizer vocab size: {tokenizer.vocab_size}\")\n    \n    # Initialize model\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    # Load data\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    # Create datasets\n    train_dataset = BrainToTextDataset(train_files, tokenizer, mode='train', augment=True)\n    val_dataset = BrainToTextDataset(val_files, tokenizer, mode='val', augment=False)\n    \n    # Create dataloaders\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    # Create trainer\n    trainer = Trainer(model, tokenizer, config)\n    \n    # Train\n    trainer.train(train_loader, val_loader)\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-01T16:18:47.088537Z","iopub.execute_input":"2025-12-01T16:18:47.089349Z","iopub.status.idle":"2025-12-01T16:37:41.402368Z","shell.execute_reply.started":"2025-12-01T16:18:47.089318Z","shell.execute_reply":"2025-12-01T16:37:41.401129Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\n\n# Compress /kaggle/working/checkpoints into /kaggle/working/checkpoints.zip\nshutil.make_archive('/kaggle/working/checkpoints', 'zip', '/kaggle/working/checkpoints')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-01T16:14:28.564834Z","iopub.execute_input":"2025-12-01T16:14:28.56573Z","iopub.status.idle":"2025-12-01T16:15:04.862506Z","shell.execute_reply.started":"2025-12-01T16:14:28.565688Z","shell.execute_reply":"2025-12-01T16:15:04.861581Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cd checkpoints","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-01T16:16:30.800735Z","iopub.execute_input":"2025-12-01T16:16:30.801361Z","iopub.status.idle":"2025-12-01T16:16:30.977233Z","shell.execute_reply.started":"2025-12-01T16:16:30.801338Z","shell.execute_reply":"2025-12-01T16:16:30.976543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cd checkpoints","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-01T16:16:54.867745Z","iopub.execute_input":"2025-12-01T16:16:54.868048Z","iopub.status.idle":"2025-12-01T16:16:55.046064Z","shell.execute_reply.started":"2025-12-01T16:16:54.868019Z","shell.execute_reply":"2025-12-01T16:16:55.044997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nDSD-NLA Inference - Generate submission.csv\nTRUE END-TO-END: Neural → Text (NO phonemes!)\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport h5py\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport glob\nimport string\nimport re\nimport torch\nfrom transformers import GPT2LMHeadModel, GPT2TokenizerFast\nfrom spellchecker import SpellChecker\n\n\n# ============================================================================\n# TEXT NORMALIZER — Strict WER Format (No punctuation)\n# ============================================================================\n\ndef normalize_for_eval(text):\n    \"\"\"\n    Remove punctuation (except apostrophe)\n    Convert to lowercase\n    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# 3. CLEANUP FILTER — reduce WER\n# ============================================================================\n\ndef clean_generated_text(text):\n    \"\"\"\n    Heuristics to improve WER without touching model:\n    - fix repeated characters\n    - fix extra spaces\n    - restore apostrophe in contractions\n    - heuristic word boundary fixes\n    \"\"\"\n    text = text.lower()\n\n    # Fix character repetitions (e.g. \"theeee\" -> \"the\")\n    text = re.sub(r\"(.)\\1{2,}\", r\"\\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# POST-PROCESSING PIPELINE\n# ============================================================================\n\nprint(\"Loading post-processing tools (GPT-2, SpellChecker, GrammarTool)...\")\n# 1. Load pre-trained GPT-2 for rescoring\ngpt2_tokenizer = GPT2TokenizerFast.from_pretrained(\"gpt2\")\ngpt2_model = GPT2LMHeadModel.from_pretrained(\"gpt2\").cuda()\ngpt2_model.eval()\n\n# 2. Load spell and grammar checking tools\nspell = SpellChecker()\nprint(\"Tools loaded successfully.\")\n\ndef spell_fix(text):\n    \"\"\"Corrects spelling and grammar mistakes in a given text.\"\"\"\n    try:\n        words = text.split()\n        # Find unknown words and get their corrections\n        unknown_words = spell.unknown(words)\n        corrected_words = [spell.correction(word) if word in unknown_words else word for word in words]\n        \n        # Join corrected words, filtering out potential None results from the spellchecker\n        text = ' '.join(filter(None, corrected_words))\n        \n    except Exception as e:\n        print(f\"Warning: Could not apply spell/grammar fix due to error: {e}\")\n    return text\n\ndef compute_gpt2_perplexity(text):\n    \"\"\"Calculates a pseudo-perplexity score for a text using GPT-2's loss.\"\"\"\n    if not text: # Handle empty strings\n        return float('inf')\n    try:\n        inputs = gpt2_tokenizer(text, return_tensors=\"pt\").to('cuda')\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\ndef select_best_candidate_by_lm(candidates):\n    \"\"\"Selects the best text from a list of candidates using GPT-2 perplexity.\"\"\"\n    best_score = float('inf')\n    best_text = candidates[0] if candidates else \"\"\n    \n    for text in candidates:\n        text_clean = text.strip()\n        # The text with the lowest perplexity (most natural) is chosen\n        ppl = compute_gpt2_perplexity(text_clean)\n        \n        if ppl < best_score:\n            best_score = ppl\n            best_text = text_clean\n            \n    return best_text\n    \n# ============================================================================\n# TOKENIZER\n# ============================================================================\n\nclass CharTokenizer:\n    \"\"\"Character-level tokenizer\"\"\"\n    def __init__(self):\n        self.pad_id = 0\n        self.bos_id = 1\n        self.eos_id = 2\n        \n        self.chars = ['<PAD>', '<BOS>', '<EOS>']\n        self.chars += list(string.ascii_lowercase)\n        self.chars += [' ']\n        self.chars += list(\"'.,!?-\")\n        \n        self.char2id = {c: i for i, c in enumerate(self.chars)}\n        self.id2char = {i: c for i, c in enumerate(self.chars)}\n        self.vocab_size = len(self.chars)\n    \n    def decode(self, ids):\n        \"\"\"Token IDs → text\"\"\"\n        if isinstance(ids, torch.Tensor):\n            ids = ids.cpu().numpy()\n        \n        chars = []\n        for i in ids:\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__(self, model_path, device='cuda', max_len=200, temperature=0.8):\n        self.device = device\n        self.max_len = max_len\n        self.temperature = temperature\n        \n        # Load tokenizer\n        self.tokenizer = CharTokenizer()\n        \n        # Load model\n        print(f\"Loading model from {model_path}...\")\n        checkpoint = torch.load(model_path, map_location=device,weights_only=False)\n        \n        config = checkpoint.get('config', {})\n        \n        # Create model \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):\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, batch_size=16):\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        predictions = []\n        \n        # Pad batch\n        from torch.nn.utils.rnn import pad_sequence\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\n            # Generate for each in batch\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__(self, model, tokenizer, beam_width=5, max_len=80):\n        self.model = model\n        self.tokenizer = tokenizer\n        self.beam_width = beam_width\n        self.max_len = max_len\n    \n    @torch.no_grad()\n    def generate(self, z_neural):\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        B = 1\n        \n        # Initialize beams\n        sequences = [[self.tokenizer.bos_id]]  # Start with BOS\n        scores = [0.0]\n        \n        for step in range(self.max_len):\n            all_candidates = []\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, 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                    new_seq = seq + [idx.item()]\n                    new_score = score + prob.item()\n                    all_candidates.append((new_seq, new_score))\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\n    @torch.no_grad()\n    def generate_candidates(self, z_neural):\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: A list of strings representing the top beam_width candidates.\n        \"\"\"\n        device = z_neural.device\n        \n        # Initialize beams\n        sequences = [[self.tokenizer.bos_id]]  # Start with BOS\n        scores = [0.0]\n        \n        for step in range(self.max_len):\n            all_candidates = []\n            \n            # This flag will be true when all beams have ended in an EOS token\n            all_beams_ended = True\n            \n            for seq, score in zip(sequences, scores):\n                if seq[-1] == self.tokenizer.eos_id:\n                    # This sequence has finished, just add it to the candidates\n                    all_candidates.append((seq, score))\n                    continue\n                \n                # If we are here, it means at least one beam is still active\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, 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                    new_seq = seq + [idx.item()]\n                    new_score = score + prob.item()\n                    all_candidates.append((new_seq, new_score))\n            \n            # If all beams have ended, we can stop early\n            if all_beams_ended:\n                break\n            \n            # Select top beam_width candidates for the next step\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        # After the loop, `sequences` holds the final top candidates.\n        # Decode all of them into text.\n        candidate_texts = [self.tokenizer.decode(seq) for seq in sequences]\n        \n        return candidate_texts\n\n# ============================================================================\n# SUBMISSION GENERATOR\n# ============================================================================\n\ndef generate_submission(\n    model_path,\n    test_data_dir,\n    output_path='submission.csv',\n    device='cuda',\n    use_beam_search=False,\n    beam_width=5\n):\n    # Initialize your custom inference engine\n    engine = InferenceEngine(model_path=model_path, device=device)\n    \n    # Initialize the beam search generator if requested\n    beam_generator = None\n    if use_beam_search:\n        beam_generator = BeamSearchGenerator(engine.model, engine.tokenizer, beam_width=beam_width)\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 = []\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(keys, desc=f\"Trials in {Path(file_path).parent.name}\", leave=False):\n                neural_features = f[key]['input_features'][:]\n                \n                final_text = \"\"\n                if use_beam_search and beam_generator:\n                    # 1. Get ALL candidates from the 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 to select the most fluent one\n                    best_text_from_lm = select_best_candidate_by_lm(beam_candidates)\n                    \n                    # 3. Apply spell and grammar correction to the selected text\n                    final_text = spell_fix(best_text_from_lm)\n                    \n                else: #greedy decoding\n                    raw_text = engine.predict(neural_features)\n                    # Apply only spell and grammar correction\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# MAIN\n# ============================================================================\n\nmodel_path = \"/kaggle/working/checkpoints/best_dsdnla_model.pt\"  \ntest_data_dir = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"              \noutput_path = \"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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-01T16:37:49.638026Z","iopub.execute_input":"2025-12-01T16:37:49.638754Z","iopub.status.idle":"2025-12-01T17:30:16.509052Z","shell.execute_reply.started":"2025-12-01T16:37:49.638723Z","shell.execute_reply":"2025-12-01T17:30:16.507975Z"}},"outputs":[],"execution_count":null}]}