{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"},{"sourceId":14334004,"sourceType":"datasetVersion","datasetId":9151343},{"sourceId":14334031,"sourceType":"datasetVersion","datasetId":9150985},{"sourceId":14342135,"sourceType":"datasetVersion","datasetId":9157364},{"sourceId":14345922,"sourceType":"datasetVersion","datasetId":9151207},{"sourceId":363134,"sourceType":"modelInstanceVersion","modelInstanceId":301514,"modelId":322000}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":12487.507943,"end_time":"2025-12-21T15:37:27.742127","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-12-21T12:09:20.234184","version":"2.6.0"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"2c8e845768a34e5685fe3199d0dd5ac3":{"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}},"3b58d3b10167438e81d76909fde8c942":{"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}},"7d799cb748d142b3a8745c62fdfe7a57":{"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_3b58d3b10167438e81d76909fde8c942","placeholder":"​","style":"IPY_MODEL_2c8e845768a34e5685fe3199d0dd5ac3","tabbable":null,"tooltip":null,"value":"Loading checkpoint shards: 100%"}},"8c3ddacef0bd4cd4b91b478c704dec70":{"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_9b700f04a36f4be1b15453e4edf9e7b3","placeholder":"​","style":"IPY_MODEL_d0dccd3b37fe4ef3a482c40fdbee9e26","tabbable":null,"tooltip":null,"value":" 3/3 [00:53&lt;00:00, 14.67s/it]"}},"9a572079c59a4e0a848e7fc36de40025":{"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":""}},"9b700f04a36f4be1b15453e4edf9e7b3":{"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}},"a5084065565543a0995f5823c3149ece":{"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_f6bb7acc3eca4646af83d2933fcb42d5","max":3,"min":0,"orientation":"horizontal","style":"IPY_MODEL_9a572079c59a4e0a848e7fc36de40025","tabbable":null,"tooltip":null,"value":3}},"b5b1287215f7414391fdca648601c3b9":{"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}},"d0dccd3b37fe4ef3a482c40fdbee9e26":{"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}},"eb126ca23c9646acbdbfeac513378a30":{"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_7d799cb748d142b3a8745c62fdfe7a57","IPY_MODEL_a5084065565543a0995f5823c3149ece","IPY_MODEL_8c3ddacef0bd4cd4b91b478c704dec70"],"layout":"IPY_MODEL_b5b1287215f7414391fdca648601c3b9","tabbable":null,"tooltip":null}},"f6bb7acc3eca4646af83d2933fcb42d5":{"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}}},"version_major":2,"version_minor":0}}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install /kaggle/input/trans-4-57/*.whl --no-deps\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:41:56.793279Z","iopub.execute_input":"2025-12-31T09:41:56.793939Z","iopub.status.idle":"2025-12-31T09:41:58.349580Z","shell.execute_reply.started":"2025-12-31T09:41:56.793892Z","shell.execute_reply":"2025-12-31T09:41:58.348892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install /kaggle/input/sympspell/*.whl --no-deps\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:41:58.351036Z","iopub.execute_input":"2025-12-31T09:41:58.351342Z","iopub.status.idle":"2025-12-31T09:41:59.585504Z","shell.execute_reply.started":"2025-12-31T09:41:58.351313Z","shell.execute_reply":"2025-12-31T09:41:59.584772Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from symspellpy import SymSpell, Verbosity\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:41:59.586732Z","iopub.execute_input":"2025-12-31T09:41:59.586970Z","iopub.status.idle":"2025-12-31T09:41:59.597357Z","shell.execute_reply.started":"2025-12-31T09:41:59.586943Z","shell.execute_reply":"2025-12-31T09:41:59.596807Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nimport math\nimport os\nimport re\nimport string\nimport unicodedata\nfrom dataclasses import dataclass\nfrom glob import glob\nfrom pathlib import Path\nfrom typing import Any, Dict, List, Optional, Pattern, Tuple, Union\n\nimport google.protobuf.message_factory\nimport h5py\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchaudio.transforms as T\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nfrom transformers import AutoModelForCausalLM, AutoTokenizer\nimport re\nimport unicodedata\nfrom typing import Dict\nfrom torchmetrics.text import WordErrorRate\nimport gc\nimport random\nfrom dataclasses import dataclass, field\n\n# Suppress TensorFlow/XLA warnings\nos.environ[\"TF_CPP_MIN_LOG_LEVEL\"] = \"3\"\nos.environ[\"TF_ENABLE_ONEDNN_OPTS\"] = \"0\"\n\n# Transformers for Rescoring\n\n# --- PASTE THE PATCH HERE ---\n\n\ndef GetPrototype(self, descriptor):\n    return self.GetMessageClass(descriptor)\n\n\nif not hasattr(google.protobuf.message_factory.MessageFactory, \"GetPrototype\"):\n    google.protobuf.message_factory.MessageFactory.GetPrototype = GetPrototype","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:41:59.599612Z","iopub.execute_input":"2025-12-31T09:41:59.599864Z","iopub.status.idle":"2025-12-31T09:42:05.154858Z","shell.execute_reply.started":"2025-12-31T09:41:59.599819Z","shell.execute_reply":"2025-12-31T09:42:05.154243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclass\nclass Config:\n    # Data\n    DATA_DIR: str = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\n    MODEL_PATH: str = \"/kaggle/input/b2t-check/b2t/checkpoints\"\n    TEST_DIR: str = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\n    OUTPUT_CSV: str = \"/kaggle/working/submission.csv\"\n    LM_MODEL_ID: str = \"/kaggle/input/qwen-3/transformers/4b/1\"\n\n    # Model\n    n_channels: int = 512\n    d_model: int = 512\n    n_encoder_layers: int = 8\n    n_decoder_layers: int = 8\n    n_heads_encoder: int = 8\n    n_heads_decoder: int = 8\n    dropout_encoder: float = 0.07  # 0.1\n    dropout_decoder: float = 0.12  # 0.2\n    # --- NEW: Define stages via lists ---\n\n    # Training\n    batch_size: int = 32  # 24\n    num_epochs: int = 100\n    learning_rate: float = 1e-3\n    weight_decay: float = 0.06\n    gradient_clip: float = 1.0\n\n    patience: int = 20\n\n    # Device\n    gpu_id: int = 0\n    num_workers: int = 16\n\n    # Logging\n    log_every: int = 100\n\n    def __post_init__(self):\n        self.device = f\"cuda:{self.gpu_id}\" if torch.cuda.is_available(\n        ) else \"cpu\"\n\n\nconfig = Config()\nprint(vars(config))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:42:05.155820Z","iopub.execute_input":"2025-12-31T09:42:05.156399Z","iopub.status.idle":"2025-12-31T09:42:05.187904Z","shell.execute_reply.started":"2025-12-31T09:42:05.156364Z","shell.execute_reply":"2025-12-31T09:42:05.187131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n\ndef normalize_for_eval(text: str) -> str:\n    \"\"\"\n    Standardizes text for ASR evaluation (applies to both Ground Truth & Hypothesis).\n\n    Key improvements:\n    1. Unicode NFKC normalization for character consistency\n    2. Proper hyphen handling (splits hyphenated words)\n    3. Robust apostrophe and contraction handling\n    4. Number normalization for consistent digit representation\n    5. Comprehensive punctuation removal with edge case handling\n\n    This function ensures that formatting differences don't artificially inflate WER.\n    \"\"\"\n    if not text or not text.strip():\n        return \"\"\n\n    # 1. Unicode Normalization (NFKC) - fixes invisible character issues\n    text = unicodedata.normalize(\"NFKC\", text)\n\n    # 2. Convert to lowercase\n    text = text.lower()\n\n    # 3. Handle hyphens first (standard ASR practice: \"wi-fi\" -> \"wi fi\")\n    text = text.replace(\"-\", \" \")\n\n    # 4. Handle smart quotes and apostrophes consistently\n    text = re.sub(r\"[`‘’‛′＇]\", \"'\", text)  # Normalize all apostrophe variants\n\n    # 5. Handle common contractions before removing punctuation\n    # Common contractions mapping (expanded form → contracted form)\n    contractions: Dict[str, str] = {\n        # basic negatives and auxiliaries\n        \"cannot\": \"can't\",\n        \"cant\": \"can't\",\n        \"do not\": \"don't\",\n        \"dont\": \"don't\",\n        \"does not\": \"doesn't\",\n        \"doesnt\": \"doesn't\",\n        \"did not\": \"didn't\",\n        \"didnt\": \"didn't\",\n        \"is not\": \"isn't\",\n        \"isnt\": \"isn't\",\n        \"are not\": \"aren't\",\n        \"arent\": \"aren't\",\n        \"was not\": \"wasn't\",\n        \"wasnt\": \"wasn't\",\n        \"were not\": \"weren't\",\n        \"werent\": \"weren't\",\n        \"have not\": \"haven't\",\n        \"havent\": \"haven't\",\n        \"has not\": \"hasn't\",\n        \"hasnt\": \"hasn't\",\n        \"had not\": \"hadn't\",\n        \"hadnt\": \"hadn't\",\n        \"would not\": \"wouldn't\",\n        \"wouldnt\": \"wouldn't\",\n        \"should not\": \"shouldn't\",\n        \"shouldnt\": \"shouldn't\",\n        \"could not\": \"couldn't\",\n        \"couldnt\": \"couldn't\",\n        \"must not\": \"mustn't\",\n        \"mustnt\": \"mustn't\",\n        \"need not\": \"needn't\",\n        \"neednt\": \"needn't\",\n        \"dare not\": \"daren't\",\n        \"darent\": \"daren't\",\n        \"might not\": \"mightn't\",\n        \"mightnt\": \"mightn't\",\n        \"ought not\": \"oughtn't\",\n        \"oughtnt\": \"oughtn't\",\n\n        # pronouns + be/have/will/modal contractions\n        \"i am\": \"i'm\",\n        \"im\": \"i'm\",\n        \"you are\": \"you're\",\n        \"youre\": \"you're\",\n        \"he is\": \"he's\",\n        \"hes\": \"he's\",\n        \"she is\": \"she's\",\n        \"shes\": \"she's\",\n        \"it is\": \"it's\",\n        \"its\": \"it's\",\n        \"we are\": \"we're\",\n        \"were\": \"we're\",          # keep for the common misspelling/token replacement\n        \"they are\": \"they're\",\n        \"theyre\": \"they're\",\n        \"i have\": \"i've\",\n        \"ive\": \"i've\",\n        \"you have\": \"you've\",\n        \"youve\": \"you've\",\n        \"we have\": \"we've\",\n        \"weve\": \"we've\",\n        \"they have\": \"they've\",\n        \"theyve\": \"they've\",\n        \"i will\": \"i'll\",\n        \"ill\": \"i'll\",\n        \"you will\": \"you'll\",\n        \"youll\": \"you'll\",\n        \"he will\": \"he'll\",\n        \"hell\": \"he'll\",\n        \"she will\": \"she'll\",\n        \"shell\": \"she'll\",\n        \"we will\": \"we'll\",\n        \"well\": \"we'll\",\n        \"they will\": \"they'll\",\n        \"theyll\": \"they'll\",\n        \"i would\": \"i'd\",\n        \"id\": \"i'd\",\n        \"you would\": \"you'd\",\n        \"youd\": \"you'd\",\n        \"he would\": \"he'd\",\n        \"hed\": \"he'd\",\n        \"she would\": \"she'd\",\n        \"shed\": \"she'd\",\n        \"we would\": \"we'd\",\n        \"wed\": \"we'd\",\n        \"they would\": \"they'd\",\n        \"theyd\": \"they'd\",\n\n        # modals + have (contractions with 've)\n        \"would have\": \"would've\",\n        \"wouldve\": \"would've\",\n        \"should have\": \"should've\",\n        \"shouldve\": \"should've\",\n        \"could have\": \"could've\",\n        \"couldve\": \"could've\",\n        \"might have\": \"might've\",\n        \"mightve\": \"might've\",\n        \"must have\": \"must've\",\n        \"mustve\": \"must've\",\n        \"need have\": \"need've\",\n        \"needve\": \"need've\",\n\n        # short forms and colloquialisms\n        \"there is\": \"there's\",\n        \"theres\": \"there's\",\n        \"here is\": \"here's\",\n        \"heres\": \"here's\",\n        \"that is\": \"that's\",\n        \"thats\": \"that's\",\n        \"what is\": \"what's\",\n        \"whats\": \"what's\",\n        \"who is\": \"who's\",\n        \"whos\": \"who's\",\n        \"when is\": \"when's\",\n        \"whens\": \"when's\",\n        \"where is\": \"where's\",\n        \"wheres\": \"where's\",\n        \"whys\": \"why's\",\n        \"how is\": \"how's\",\n        \"hows\": \"how's\",\n\n        # time and other common contractions\n        \"of the clock\": \"o'clock\",\n        \"oclock\": \"o'clock\",\n        \"it is not\": \"it isn't\",\n        \"aint\": \"ain't\",\n        \"am not\": \"ain't\",\n        \"yall\": \"y'all\",\n        \"ya'll\": \"y'all\",\n        \"cannot've\": \"can't've\",\n        \"cantve\": \"can't've\",\n        \"wouldn't've\": \"wouldn't've\",\n        \"wouldntve\": \"wouldn't've\",\n\n        # less common/colloquial contractions\n        \"let us\": \"let's\",\n        \"lets\": \"let's\",\n        \"which is\": \"which's\",\n        \"whichs\": \"which's\",\n        \"that would\": \"that'd\",\n        \"thatd\": \"that'd\",\n        \"who would\": \"who'd\",\n        \"whod\": \"who'd\",\n        \"what would\": \"what'd\",\n        \"whatd\": \"what'd\",\n        \"there would\": \"there'd\",\n        \"thered\": \"there'd\",\n        \"here would\": \"here'd\",\n        \"hered\": \"here'd\",\n\n        # keep lowercased variants for robust token replacement\n        \"imma\": \"I'ma\",\n        \"lemme\": \"lemme\",   # intentionally unchanged colloquial, left as-is\n        \"gonna\": \"gonna\",\n        \"gotta\": \"gotta\"\n    }\n\n    # Apply contractions mapping with word boundaries\n    for wrong, correct in contractions.items():\n        text = re.sub(rf'\\b{wrong}\\b', correct, text)\n\n    # 6. Remove punctuation except apostrophes within words\n    text = re.sub(r'[^\\w\\s\\']', ' ', text)\n\n    # 7. Fix apostrophe spacing issues\n    text = re.sub(r\"(\\w)\\s+'\\s*(\\w)\", r\"\\1'\\2\", text)  # \"don 't\" -> \"don't\"\n    text = re.sub(r\"'\\s+(\\w)\", r\"'\\1\", text)           # \"' t\" -> \"'t\"\n    text = re.sub(r\"(\\w)\\s+'\", r\"\\1'\", text)           # \"don ' \" -> \"don'\"\n\n    # 8. Handle apostrophes at word boundaries\n    # Remove standalone apostrophes\n    text = re.sub(r\"\\s+'\\s+\", \" \", text)\n    # Remove apostrophes at start/end\n    text = re.sub(r\"^'|\\s+'$|'$\", \"\", text)\n\n    # 10. Normalize whitespace and trim\n    text = re.sub(r'\\s+', ' ', text).strip()\n\n    return text\n\n\n\ndef clean_generated_text(text: str) -> str:\n    \"\"\"\n    Heuristics to repair model artifacts (Hypothesis only).\n\n    Key improvements:\n    1. Word-level repetition cleanup for ASR stutter patterns\n    2. Character-level repetition with vowel/consonant differentiation\n    3. Contraction spacing repair\n    4. **NEW: Remove word repetitions at the end** (common ASR artifact)\n    5. Whitespace normalization\n\n    This function fixes common ASR model artifacts without \"cheating\" WER by correcting\n    actual content errors (spelling mistakes remain intact for fair evaluation).\n    \"\"\"\n    if not text or not text.strip():\n        return \"\"\n\n    # 1. Basic lowercasing and whitespace normalization\n    text = text.lower()\n    text = \" \".join(text.split())  # Normalize all whitespace\n\n    # 2. Fix contraction spacing patterns first\n    text = re.sub(r\"(\\w)\\s+'\\s*(\\w)\", r\"\\1'\\2\", text)  # \"don 't\" -> \"don't\"\n    text = re.sub(r\"(\\w)\\s+'(\\s|$)\", r\"\\1'\\2\", text)   # \"don ' \" -> \"don' \"\n    # Remove apostrophes at start/end\n    text = re.sub(r\"^'\\s+|\\s+'$\", \"\", text)\n\n    # 3. Word-level repetition cleanup (ASR stutter patterns)\n    # Conservative approach: only target common function words and obvious stutters\n\n    # First: collapse 3+ identical words to 2 (very conservative)\n    text = re.sub(r'\\b(\\w+)\\s+(\\1\\s+){2,}', r'\\1 \\1 ', text)\n\n    # Second: collapse common function word repetitions\n    common_repeat_words = r'(the|and|that|this|is|are|was|were|be|to|of|in|on|at|for|with|but|or)'\n    text = re.sub(rf'\\b{common_repeat_words}\\s+\\1\\b', r'\\1', text)\n\n    # Third: handle pronoun repetitions\n    text = re.sub(r'\\b(i|i\\'m|i\\'ve|i\\'ll|i\\'d)\\s+\\1\\b',\n                  r'\\1', text)  # \"i i\" -> \"i\"\n    text = re.sub(r'\\b(you|you\\'re|you\\'ve|you\\'ll|you\\'d)\\s+\\1\\b',\n                  r'\\1', text)  # \"you you\" -> \"you\"\n\n    # 4. **NEW: Remove word repetitions at the end of text**\n    # This handles common ASR artifacts like \"hello world world\" -> \"hello world\"\n    # or \"this is the end end end\" -> \"this is the end\"\n\n    # Pattern 1: Single word repetition at end (e.g., \"word word\" at end)\n    text = re.sub(r'\\b(\\w+)\\s+\\1$', r'\\1', text)\n\n    # Pattern 2: Multiple word repetitions at end (e.g., \"word word word\" at end)\n    # This handles 2+ repetitions of the same word at the end\n    text = re.sub(r'\\b(\\w+)(?:\\s+\\1)+$', r'\\1', text)\n\n    # Pattern 3: Handle cases with punctuation at the end\n    text = re.sub(r'\\b(\\w+)(?:\\s+\\1)+[.!?]*$', r'\\1', text)\n\n    # 5. Character-level repetition cleanup (model stuttering)\n    # Vowels: reduce 3+ repetitions to 2 (preserve legitimate doubles like \"book\")\n    text = re.sub(r\"([aeiou])\\1{2,}\", r\"\\1\\1\", text)\n\n    # Consonants: reduce 3+ repetitions to 1 (English has few legitimate triple consonants)\n    text = re.sub(r\"([bcdfghjklmnpqrstvwxyz])\\1{2,}\", r\"\\1\", text)\n\n    # 6. Remove standalone/dangling apostrophes\n    text = re.sub(r\" ' \", \" \", text)\n    text = re.sub(r\"\\s+'\\s*\", \" \", text)\n\n    # 7. Final whitespace cleanup\n    text = re.sub(r'\\s+', ' ', text).strip()\n\n    return text\n\n\n# Usage example:\n# clean_text = clean_generated_text(best_text)  # Hypothesis-specific cleaning\n# final_text = normalize_for_eval(clean_text)   # Standard evaluation normalization","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:42:05.189225Z","iopub.execute_input":"2025-12-31T09:42:05.189760Z","iopub.status.idle":"2025-12-31T09:42:05.211465Z","shell.execute_reply.started":"2025-12-31T09:42:05.189735Z","shell.execute_reply":"2025-12-31T09:42:05.210748Z"}},"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 re\nimport string\nfrom collections import Counter\nfrom typing import Dict, List, Optional, Union\n\nimport torch\nimport torch.nn as nn\n# ============================================================================#\n# 0. MODEL DEFINITION (must match inference script)                           #\n# ============================================================================#\nimport torch.nn.functional as F  # <-- This is the F import\nfrom torchaudio.models import Conformer\n\nfrom transformers import Wav2Vec2ConformerConfig, Wav2Vec2ConformerModel\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.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 LearnedPositionalEncoding(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        # Learnable parameter instead of fixed sinusoids\n        self.pos_embedding = nn.Parameter(\n            torch.randn(1, max_len, d_model) * 0.02)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        seq_len = x.size(1)\n        x = x + self.pos_embedding[:, :seq_len, :]\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\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=4, padding=2),\n            nn.BatchNorm1d(d_model),\n            nn.GELU(),\n            nn.Dropout(dropout),\n        )\n\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        #     nn.Conv1d(d_model, 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        # 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, norm=nn.LayerNorm(d_model))\n\n        # --- Conformer Encoder ---\n        self.conformer = Conformer(\n            input_dim=d_model,\n            num_heads=n_heads,\n            ffn_dim=d_model * 1,\n            num_layers=n_layers,\n            depthwise_conv_kernel_size=21,  # Typical kernel size for Conformer\n            dropout=dropout,\n            use_group_norm=True,\n            # convolution_first=True\n        )\n\n        # # Wav2Vec 2.0 style Conformer configuration\n        # config = Wav2Vec2ConformerConfig(\n        #     hidden_size=d_model,\n        #     num_attention_heads=n_heads,\n        #     intermediate_size=d_model * 4,\n        #     num_hidden_layers=n_layers,\n        #     conv_depthwise_kernel_size=31,  # Wav2Vec 2.0 typically uses larger kernels\n        #     hidden_dropout=dropout,\n        #     attention_dropout=dropout,\n        #     feat_proj_dropout=dropout,\n        #     layer_norm_eps=1e-5\n        # )\n\n        # self.conformer = Wav2Vec2ConformerModel(config)\n\n        self.layer_norm = nn.LayerNorm(d_model)\n\n    def _update_mask(self, mask: torch.Tensor, stride: int, kernel_size: int, target_len: int) -> torch.Tensor:\n        \"\"\"\n        Updates mask for a SINGLE stage using MaxPool logic.\n        \"\"\"\n        if mask is None:\n            return None\n\n        # Invert: True(Pad) -> 0.0, False(Valid) -> 1.0\n        # (We want 1.0 if ANY data is valid)\n        x = (~mask).float().unsqueeze(1)  # (B, 1, T)\n\n        padding = kernel_size // 2\n\n        # MaxPool mirrors the conv geometry\n        x = F.max_pool1d(\n            x, kernel_size=kernel_size, stride=stride, padding=padding, ceil_mode=False\n        )\n\n        # Resize to match feature map exactly (handles edge cases)\n        if x.size(2) != target_len:\n            x = F.interpolate(x, size=target_len, mode='nearest')\n\n        # Invert back: 1.0 -> False, 0.0 -> True\n        return ~(x.bool().squeeze(1))\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[:, ::4]\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        if mask is not None:\n            mask = self._update_mask(\n                mask,\n                stride=4,\n                kernel_size=5,\n                target_len=x.size(1)\n            )\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        # --- Conformer expects lengths instead of boolean mask ---\n        if mask is not None:\n            # Convert boolean mask (True=padding) to lengths\n            # (~mask) gives True for valid tokens, sum gives count of valid tokens\n            input_lengths = (~mask).sum(dim=1).long()\n        else:\n            # If no mask, all sequences are full length\n            input_lengths = torch.full(\n                (x.size(0),),\n                x.size(1),\n                dtype=torch.long,\n                device=x.device\n            )\n\n        # --- Run through Conformer ---\n        # torchaudio Conformer returns (output, output_lengths)\n        # ctc length!!\n        x, output_lengths = self.conformer(x, input_lengths)\n\n        # --- Update mask based on output lengths ---\n        if mask is not None:\n            # Create new mask based on output lengths\n            max_len = x.size(1)\n            # Create range tensor [0, 1, 2, ..., max_len-1] for each batch item\n            range_tensor = torch.arange(\n                max_len, device=x.device).expand(x.size(0), max_len)\n            # New mask: True where position >= length (padding positions)\n            new_mask = range_tensor >= output_lengths.unsqueeze(1)\n            mask = new_mask\n        # else: keep mask as None\n\n        return x, mask\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\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 * 1,\n            dropout=dropout,\n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=True,\n        )\n        self.decoder = nn.TransformerDecoder(\n            decoder_layer, num_layers=n_layers, norm=nn.LayerNorm(d_model))\n        # self.dropout_classifier = nn.Dropout(dropout_classifier)\n        # self.final_norm = nn.LayerNorm(d_model)  # Add this\n        # self.output_head = nn.Sequential(\n        #     nn.Linear(d_model, d_model // 2),\n        #     nn.GELU(),\n        #     nn.Dropout(dropout_classifier),\n        #     nn.Linear(d_model // 2, vocab_size)\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        mask=None,\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(\n                tgt_emb, z_neural, tgt_mask=causal_mask, memory_key_padding_mask=mask)\n            # out = self.dropout_classifier(out)\n            # out = self.final_norm(out)  # Apply final layer norm\n            out = self.output_proj(out)\n            # logits = self.output_head(out)  # Use output head for logits\n            return out\n\n        # For training we always pass target_tokens, inference is handled by DSDNLA.inference\n        raise RuntimeError(\n            \"TextDecoder.forward called without target_tokens in training mode\"\n        )\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\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_encoder: int = 8,\n            n_heads_decoder: int = 8,\n            dropout_encoder: float = 0.1,\n            dropout_decoder: float = 0.1):\n        super().__init__()\n        self.tokenizer = CharTokenizer()\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_encoder,\n            dropout=dropout_encoder,\n        )\n        # freeze the encoder !!\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_decoder,\n            dropout=dropout_decoder\n        )\n        self.backbone_frozen = False\n\n    def freeze(self):\n        \"\"\"\n        Freezes the Feature Extractors and the Main Transformer Decoder.\n        Keeps the Pooling mechanism and Output MLP trainable.\n        \"\"\"\n        print(\"❄️ Freezing Backbone (Projections + Transformer Decoder)...\")\n        # 1. Freeze Dynamic Projection\n        for param in self.neural_encoder.parameters():\n            param.requires_grad = False\n        self.backbone_frozen = True\n\n    def train(self, mode=True):\n        \"\"\"\n        Custom train mode.\n        If the backbone is frozen, we force it to remain in eval mode\n        even when the rest of the model is set to train.\n        \"\"\"\n        # 1. Set the whole model to the requested mode (usually True)\n        super().train(mode)\n\n        # # 2. If backbone is frozen, force neural encoder to eval mode\n        if self.backbone_frozen:\n            self.neural_encoder.eval()\n        else:\n            # Only set to train if not frozen and mode is True\n            self.neural_encoder.train(mode)\n\n        # 3. Always keep decoder in the requested mode\n        self.text_decoder.train(mode)\n\n        return self\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, mask = self.neural_encoder(neural_features, mask=neural_mask)\n        logits = self.text_decoder(\n            z_neural, target_tokens=target_tokens, mask=mask)\n        return logits\n\n    @torch.inference_mode()\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((B, 1), bos_id, dtype=torch.long, device=device)\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(\n                T).to(device)\n            out = self.text_decoder.decoder(\n                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    @torch.inference_mode()\n    def generate_candidate(self, neural_features: torch.Tensor,  neural_mask: Optional[torch.Tensor] = None) -> List[str]:\n        \"\"\"\n        Returns list of top-K decoded strings\n        z_neural: (1, L, d_model)\n        \"\"\"\n        self.max_len = 100\n        self.beam_width = 10\n\n        self.neural_encoder.eval()\n        self.text_decoder.eval()\n\n        z_neural = self.neural_encoder(neural_features, mask=neural_mask)\n\n        device = z_neural.device\n        # Beams: (sequence_list, score)\n        beams = [([self.tokenizer.bos_id], 0.0)]\n\n        for _ in range(self.max_len):\n            candidates = []\n            all_ended = True\n\n            for seq, score in beams:\n                if seq[-1] == self.tokenizer.eos_id:\n                    candidates.append((seq, score))\n                    continue\n\n                all_ended = False\n\n                # Expand\n                seq_tensor = torch.tensor(\n                    [seq], dtype=torch.long, device=device)\n                logits = self.text_decoder(\n                    z_neural, target_tokens=seq_tensor)\n                next_logits = logits[0, -1, :]\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 lp, idx in zip(topk_probs, topk_ids):\n                    idx = idx.item()\n                    lp = lp.item()\n\n                    # Length Penalty\n                    new_len = len(seq) + 1\n                    # Update score (additive log prob)\n                    # We normalize strictly by length^0.7\n                    prev_raw_score = score * \\\n                        (len(seq) ** 0.7) if len(seq) > 0 else 0\n                    new_raw_score = prev_raw_score + lp\n                    new_score = new_raw_score / (new_len**0.7)\n\n                    candidates.append((seq + [idx], new_score))\n\n            if all_ended or not candidates:\n                break\n\n            # Prune to beam width\n            candidates.sort(key=lambda x: x[1], reverse=True)\n            beams = candidates[: self.beam_width]\n\n        # Stop if all top beams are ended\n        # Decode all beams\n        best_text = beams[0][0]\n        best_text = self.tokenizer.decode(best_text)\n        clean_text = clean_generated_text(best_text)\n        final_text = normalize_for_eval(clean_text)\n\n        return final_text\n\n\n# ============================================================================#\n# 1. TOKENIZER - Character-level                                             #\n# ============================================================================#\n\nclass CharTokenizer:\n    \"\"\"\n    Character-level tokenizer that only considers a-z characters and space.\n    Special tokens:\n    0: PAD - padding token\n    1: BOS - beginning of sequence\n    2: EOS - end of sequence\n    3-28: a-z characters (26 letters)\n    29: space character (replaces UNK)\n    \"\"\"\n\n    def __init__(self,\n                 preserve_case: bool = False,\n                 normalize_unicode: bool = True):\n        self.pad_id = 0\n        self.bos_id = 1\n        self.eos_id = 2\n\n        # Base character set with only a-z and space\n        base_chars = [\"<PAD>\", \"<BOS>\", \"<EOS>\"]\n        # Add lowercase letters a-z\n        base_chars += list(string.ascii_lowercase)\n        # Add space as the last character class (replaces UNK)\n        base_chars.append(\" \")\n\n        self.chars = base_chars\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        # Space ID is the last token (replaces UNK functionality)\n        self.space_id = self.char2id[\" \"]\n\n        self.preserve_case = preserve_case\n        self.normalize_unicode = normalize_unicode\n\n        # Clean regex: only allow a-z and whitespace (no numbers, punctuation, etc.)\n        self.clean_regex = re.compile(r\"[^a-z\\s]\")\n\n    def _normalize_text(self, text: str) -> str:\n        \"\"\"Normalize text by handling case and keeping only a-z characters and spaces\"\"\"\n        if not self.preserve_case:\n            text = text.lower()\n\n        if self.normalize_unicode:\n            # Replace common Unicode variants with space or remove\n            text = text.replace('\\u2019', \" \")  # Right single quote → space\n            text = text.replace('\\u2018', \" \")  # Left single quote → space\n            text = text.replace('\\u201c', \" \")  # Left double quote → space\n            text = text.replace('\\u201d', \" \")  # Right double quote → space\n            text = text.replace('\\u2013', \" \")  # En dash → space\n            text = text.replace('\\u2014', \" \")  # Em dash → space\n            text = text.replace('\\u00a0', \" \")  # Non-breaking space → space\n            text = text.replace('\\t', \" \")      # Tab → space\n            text = text.replace('\\n', \" \")      # Newline → space\n            text = text.replace('\\r', \" \")      # Carriage return → space\n\n        # Clean: remove anything not in a-z or whitespace\n        text = self.clean_regex.sub('', text)\n\n        # Replace multiple spaces with single space\n        text = re.sub(r'\\s+', ' ', text)\n\n        # Strip leading/trailing spaces\n        text = text.strip()\n\n        return text\n\n    def encode(self, text: str) -> torch.Tensor:\n        \"\"\"Text → token IDs (Tensor of shape [T])\"\"\"\n        text = self._normalize_text(text)\n\n        ids = [self.bos_id]\n        for c in text:\n            if c == \" \":\n                char_id = self.space_id\n            elif c in self.char2id:\n                char_id = self.char2id[c]\n            else:\n                # This shouldn't happen due to normalization, but fallback to space\n                char_id = self.space_id\n            ids.append(char_id)\n        ids.append(self.eos_id)\n\n        return torch.tensor(ids, dtype=torch.long)\n\n    def decode(self, ids: Union[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 = []\n        for i in ids_iter:\n            if i == self.eos_id:\n                break\n            if i == self.bos_id:\n                continue\n            if i == self.pad_id:\n                continue\n            if 0 <= i < self.vocab_size:\n                chars.append(self.id2char[i])\n            else:\n                # Out of vocabulary - use space (shouldn't happen with proper encoding)\n                chars.append(\" \")\n\n        return \"\".join(chars)\n\n    def get_vocab(self) -> Dict[str, int]:\n        \"\"\"Return vocabulary mapping\"\"\"\n        return self.char2id.copy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:42:05.212654Z","iopub.execute_input":"2025-12-31T09:42:05.213001Z","iopub.status.idle":"2025-12-31T09:42:08.560582Z","shell.execute_reply.started":"2025-12-31T09:42:05.212978Z","shell.execute_reply":"2025-12-31T09:42:08.560007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================#\n# 2. DATASET for Brain-to-Text                                               #\n# ============================================================================#\n\n\nfrom typing import List, Tuple, Dict, Any\nimport random\nfrom timm.utils import ModelEmaV3\nfrom typing import List, Dict, Any, Optional\nfrom tqdm import tqdm\nimport h5py\nfrom torch.utils.data import Dataset\nfrom sklearn.preprocessing import StandardScaler\nfrom typing import Tuple\nimport numpy as np\nimport torch\nimport math\nimport torch.nn.functional as F\nimport torch.nn as nn\n# Try to import scipy for Gaussian smoothing, fall back to PyTorch implementation\n\nfrom scipy.ndimage import gaussian_filter1d\nHAS_SCIPY = True\n\n\nclass PhysioAwareAugment(nn.Module):\n    def __init__(self, num_electrodes: int = 256, bin_duration_ms: int = 20):\n        super().__init__()\n        self.num_electrodes = num_electrodes\n        self.bin_duration_ms = bin_duration_ms\n\n        # Array definitions\n        self.arrays = {\n            'ventral_6v': (0, 64),\n            'area_4': (64, 128),\n            'area_55b': (128, 192),\n            'dorsal_6v': (192, 256),\n        }\n\n        self.critical_arrays = ['area_4', 'area_55b']\n\n        # Feature type indices\n        self.TC = slice(0, 256)\n        self.SBP = slice(256, 512)\n\n        # Gaussian smoothing parameters\n        self.smooth_kernel_std = 2.0\n        self.smooth_kernel_size = 100\n        self.has_scipy = HAS_SCIPY\n\n    def _ensure_batch_dim(self, x: torch.Tensor) -> Tuple[torch.Tensor, bool]:\n        \"\"\"Ensure input has batch dimension\"\"\"\n        if x.dim() == 2:  # (T, C)\n            return x.unsqueeze(0), True  # (1, T, C)\n        elif x.dim() == 3:  # (B, T, C)\n            return x, False\n        else:\n            raise ValueError(f\"Expected 2D or 3D input, got shape {x.shape}\")\n\n    def _restore_shape(self, x: torch.Tensor, squeeze: bool) -> torch.Tensor:\n        \"\"\"Restore original shape after augmentation\"\"\"\n        return x.squeeze(0) if squeeze else x\n\n    def _temporal_warp(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Clean, efficient temporal warping that actually works\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        if T < 10:  # Safety check for short sequences\n            return x\n\n        # Create warp field with 5 control points\n        control_pts = torch.linspace(0, 1, 5, device=device)\n        offsets = (torch.rand(5, device=device) - 0.5) * 0.15\n        warped_pts = torch.clamp(control_pts + offsets, 0, 1)\n        warped_pts, _ = torch.sort(warped_pts)  # Ensure monotonic\n\n        # Interpolate to get full warp field (shape: [T])\n        warp_field = F.interpolate(\n            warped_pts.view(1, 1, -1),\n            size=T,\n            mode='linear',\n            align_corners=True\n        ).squeeze(0).squeeze(0)  # Ensure shape is [T]\n\n        # Map to indices and clamp\n        indices = warp_field * (T - 1)\n        indices = torch.clamp(indices, 0, T - 1)\n\n        # Get floor and ceiling indices\n        floor_idx = torch.floor(indices).long()\n        ceil_idx = torch.ceil(indices).long()\n        ceil_idx = torch.clamp(ceil_idx, 0, T - 1)\n\n        # Interpolation weights (shape: [T])\n        weight = indices - floor_idx.float()\n\n        # Create output tensor\n        x_warped = torch.zeros_like(x)\n\n        # Vectorized interpolation per batch\n        for b in range(B):\n            # For each batch, interpolate all channels at once\n            floor_vals = x[b, floor_idx, :]  # [T, C]\n            ceil_vals = x[b, ceil_idx, :]    # [T, C]\n\n            # Expand weight to match channel dimension: [T] -> [T, 1]\n            weight_expanded = weight.view(-1, 1)\n\n            # Interpolate: (1-weight)*floor + weight*ceil\n            x_warped[b] = floor_vals * \\\n                (1 - weight_expanded) + ceil_vals * weight_expanded\n\n        return x_warped\n\n    def _gauss_smooth(self, inputs: torch.Tensor, smooth_kernel_std: float = None,\n                      smooth_kernel_size: int = None, padding: str = 'same') -> torch.Tensor:\n        \"\"\"\n        Applies 1D Gaussian smoothing with proper fallback if scipy is not available.\n\n        Args:\n            inputs (tensor): B x T x N tensor\n            smooth_kernel_std (float): Standard deviation of Gaussian kernel\n            smooth_kernel_size (int): Size of Gaussian kernel\n            padding (str): Padding mode ('same' or 'valid')\n        \"\"\"\n        if smooth_kernel_std is None:\n            smooth_kernel_std = self.smooth_kernel_std\n        if smooth_kernel_size is None:\n            smooth_kernel_size = self.smooth_kernel_size\n\n        device = inputs.device\n        B, T, C = inputs.shape\n\n        if self.has_scipy:\n            return self._gauss_smooth_scipy(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n        else:\n            return self._gauss_smooth_torch(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n\n    def _gauss_smooth_scipy(self, inputs: torch.Tensor, device: torch.device,\n                            smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                            padding: str = 'same') -> torch.Tensor:\n        \"\"\"Gaussian smoothing using scipy (more accurate)\"\"\"\n        # Get Gaussian kernel\n        inp = np.zeros(smooth_kernel_size, dtype=np.float32)\n        inp[smooth_kernel_size // 2] = 1\n        gaussKernel = gaussian_filter1d(inp, smooth_kernel_std)\n        validIdx = np.argwhere(gaussKernel > 0.01)\n        gaussKernel = gaussKernel[validIdx]\n        gaussKernel = np.squeeze(gaussKernel / np.sum(gaussKernel))\n\n        # Convert to tensor\n        gaussKernel = torch.tensor(\n            gaussKernel, dtype=torch.float32, device=device)\n        gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n        # Prepare convolution\n        B, T, C = inputs.shape\n        inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n        gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n        # Perform convolution\n        smoothed = F.conv1d(inputs_perm, gaussKernel,\n                            padding=padding, groups=C)\n        return smoothed.permute(0, 2, 1)  # [B, T, C]\n\n    def _gauss_smooth_torch(self, inputs: torch.Tensor, device: torch.device,\n                            smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                            padding: str = 'same') -> torch.Tensor:\n        \"\"\"Gaussian smoothing using pure PyTorch (fallback)\"\"\"\n        # Create Gaussian kernel using PyTorch\n        x = torch.arange(-(smooth_kernel_size//2), smooth_kernel_size //\n                         2 + 1, device=device, dtype=torch.float32)\n        gaussKernel = torch.exp(-x**2 / (2 * smooth_kernel_std**2))\n        gaussKernel = gaussKernel / gaussKernel.sum()\n\n        # Ensure kernel has odd size\n        if gaussKernel.size(0) % 2 == 0:\n            gaussKernel = gaussKernel[:-1]\n\n        gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n        # Prepare convolution\n        B, T, C = inputs.shape\n        inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n        gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n        # Handle padding\n        if padding == 'same':\n            padding_size = gaussKernel.size(2) // 2\n        else:\n            padding_size = 0\n\n        # Perform convolution\n        smoothed = F.conv1d(inputs_perm, gaussKernel,\n                            padding=padding_size, groups=C)\n        return smoothed.permute(0, 2, 1)  # [B, T, C]\n\n    def _create_coupled_electrode_mask(self, x: torch.Tensor, keep_prob: float = 0.85) -> torch.Tensor:\n        \"\"\"Create mask where both TC and SBP features for same electrode are dropped together\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        # Base mask for 256 electrodes\n        elec_mask = torch.bernoulli(torch.ones(\n            self.num_electrodes, device=device) * keep_prob)\n\n        # Add spatial correlation within arrays\n        for array_name, (start_idx, end_idx) in self.arrays.items():\n            center_elec = torch.randint(\n                start_idx, end_idx, (1,), device=device).item()\n            distances = torch.arange(\n                start_idx, end_idx, device=device).float() - center_elec\n            spatial_weights = torch.exp(-torch.abs(distances) / 10.0)\n\n            spatial_bias = 0.25 * spatial_weights\n            p = torch.clamp(keep_prob - spatial_bias, 0.05, 1.0)\n\n            array_mask = torch.bernoulli(p)\n            elec_mask[start_idx:end_idx] = array_mask\n\n        # Expand to 512 features: [TC_mask, SBP_mask]\n        full_mask = torch.cat([elec_mask, elec_mask])\n        return full_mask.view(1, 1, -1)  # (1, 1, 512)\n\n    def _create_array_mask(self, x: torch.Tensor, dropout_prob: float = 0.1) -> torch.Tensor:\n        \"\"\"Create mask that drops entire arrays with anatomical awareness\"\"\"\n        B, T, C = x.shape\n        device = x.device\n        mask = torch.ones(512, device=device)\n\n        for array_name, (start_idx, end_idx) in self.arrays.items():\n            is_critical = array_name in self.critical_arrays\n            array_dropout_prob = dropout_prob * 0.3 if is_critical else dropout_prob\n\n            if torch.rand(1, device=device) < array_dropout_prob:\n                mask[start_idx:end_idx] = 0  # TC features\n                mask[start_idx + 256:end_idx + 256] = 0  # SBP features\n\n        return mask.view(1, 1, -1)  # (1, 1, 512)\n\n    def _apply_feature_specific_noise(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply physiologically appropriate noise to different feature types\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        tc_part = x[:, :, self.TC].clone()\n        sbp_part = x[:, :, self.SBP].clone()\n\n        # Threshold Crossings (TC): multiplicative noise\n        if torch.rand(1, device=device) < 0.3:\n            scale_factor = 0.9 + torch.rand(1, device=device).item() * 0.2\n            tc_part = tc_part * scale_factor\n\n            if torch.rand(1, device=device) < 0.5:\n                noise_level = 0.1 + torch.rand(1, device=device).item() * 0.1\n                lam = torch.clamp(tc_part.abs() * noise_level, 0, 5)\n                poisson_noise = torch.poisson(lam)\n                tc_part = tc_part + poisson_noise\n\n        # Spike Band Power (SBP): colored noise\n        if torch.rand(1, device=device) < 0.4:\n            alpha = 0.85 + torch.rand(1, device=device).item() * 0.1\n            noise = torch.randn_like(sbp_part, device=device)\n            filtered_noise = torch.zeros_like(noise)\n\n            for t in range(1, T):\n                filtered_noise[:, t] = alpha * \\\n                    filtered_noise[:, t-1] + (1-alpha) * noise[:, t]\n\n            sbp_part = sbp_part + filtered_noise * \\\n                (0.03 + torch.rand(1, device=device).item() * 0.02)\n\n        return torch.cat([tc_part, sbp_part], dim=2)\n\n    def _apply_temporal_masking(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply speech-appropriate temporal masking with smooth edges\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        if torch.rand(1, device=device) < 0.25:\n            max_mask_duration_ms = 200\n            max_mask_bins = max_mask_duration_ms // self.bin_duration_ms\n            num_masks = torch.randint(1, 3, (1,), device=device).item()\n\n            for _ in range(num_masks):\n                mask_duration = torch.randint(\n                    1, max_mask_bins + 1, (1,), device=device).item()\n                start_idx = torch.randint(\n                    0, max(1, T - mask_duration), (1,), device=device).item()\n\n                mask = torch.ones(T, device=device)\n                mask[start_idx:start_idx + mask_duration] = 0\n\n                window_size = min(2, mask_duration // 2)\n                ramp = torch.linspace(0, 1, steps=window_size+1, device=device)\n\n                for i in range(window_size):\n                    if start_idx - i - 1 >= 0:\n                        mask[start_idx - i - 1] = ramp[-(i+2)]\n                    if start_idx + mask_duration + i < T:\n                        mask[start_idx + mask_duration + i] = ramp[i+1]\n\n                x = x * mask.view(1, T, 1)\n\n        return x\n\n    def _apply_global_modulation(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply slow global modulation simulating behavioral state changes\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        modulation_freq = 0.05 + torch.rand(1, device=device).item() * 0.1\n        time_vec = torch.arange(T, device=device).float() / 50.0\n\n        tc_modulation = 0.95 + 0.05 * \\\n            torch.sin(2 * np.pi * modulation_freq * time_vec +\n                      torch.rand(1, device=device).item() * 2 * np.pi)\n        sbp_modulation = 0.85 + 0.3 * \\\n            torch.sin(2 * np.pi * modulation_freq * time_vec +\n                      torch.rand(1, device=device).item() * 2 * np.pi)\n\n        x[:, :, self.TC] *= tc_modulation.view(1, T, 1)\n        x[:, :, self.SBP] *= sbp_modulation.view(1, T, 1)\n\n        return x\n\n    def _apply_slow_drift(self, x: torch.Tensor, rank: int = 2, scale: float = 0.01) -> torch.Tensor:\n        \"\"\"Apply low-rank slow drift for cross-session robustness\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        drift_t = torch.randn(B, T, rank, device=device)\n        drift_t = F.avg_pool1d(drift_t.permute(\n            0, 2, 1), kernel_size=25, stride=1, padding=12).permute(0, 2, 1)\n\n        drift_f = torch.randn(rank, C, device=device)\n        drift = torch.matmul(drift_t, drift_f)\n\n        return x + scale * torch.tanh(drift)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Physiologically-aware augmentation for speech motor cortex BCI data.\n\n        Input:\n            x: torch.Tensor of shape (T, 512) or (B, T, 512)\n\n        Output:\n            Augmented tensor of same shape\n        \"\"\"\n        x, squeeze = self._ensure_batch_dim(x)\n        device = x.device\n\n        # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n        if torch.rand(1, device=device) < 0.35:\n            x = self._temporal_warp(x)\n\n        # 2. Coupled Electrode Dropout (p=0.25) - PHYSIOLOGICALLY ACCURATE\n        if torch.rand(1, device=device) < 0.25:\n            electrode_mask = self._create_coupled_electrode_mask(\n                x, keep_prob=random.uniform(\n                    0.75, 0.95))\n            x = x * electrode_mask\n\n        # 3. Array-Level Dropout (p=0.15) - SPATIAL ROBUSTNESS\n        if torch.rand(1, device=device) < 0.15:\n            array_mask = self._create_array_mask(x, dropout_prob=random.uniform(\n                0.1, 3.0))\n            x = x * array_mask\n\n        # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n        if torch.rand(1, device=device) < 0.45:\n            x = self._apply_feature_specific_noise(x)\n\n        # 5. Temporal Masking (p=0.3) - SHORT-TERM ARTIFACTS\n        if torch.rand(1, device=device) < 0.3:\n            x = self._apply_temporal_masking(x)\n\n        # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n        if torch.rand(1, device=device) < 0.2:\n            x = self._apply_global_modulation(x)\n\n        # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n        if torch.rand(1, device=device) < 0.4:\n            x = self._apply_slow_drift(x, rank=random.randint(\n                2, 4), scale=random.uniform(0.005, 0.03))\n        # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n        if torch.rand(1, device=device) < 0.25:\n            x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n                0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n        return self._restore_shape(x, squeeze)\n\n    def forward_tta(self, x: torch.Tensor, num_augments: int = 1, seed: int = None) -> torch.Tensor:\n        \"\"\"\n        Test-Time Augmentation (TTA) method that applies mild, controlled augmentations\n        close to the original data distribution for robust inference.\n\n        Args:\n            x: Input tensor of shape (T, 512) or (B, T, 512)\n            num_augments: Number of augmented versions to generate (default: 1)\n            seed: Optional random seed for reproducibility\n\n        Returns:\n            Augmented tensor. If num_augments > 1, returns tensor of shape (num_augments, B, T, C)\n            otherwise returns same shape as input\n        \"\"\"\n        if seed is not None:\n            torch.manual_seed(seed)\n            if x.device.type == 'cuda':\n                torch.cuda.manual_seed(seed)\n\n        x_original, squeeze = self._ensure_batch_dim(x)\n        B, T, C = x_original.shape\n        device = x_original.device\n\n        # Store multiple augmented versions\n        augmented_versions = []\n\n        for i in range(num_augments):\n            x = x_original.clone()\n\n            # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n            if torch.rand(1, device=device) < 0.35:\n                x = self._temporal_warp(x)\n\n            # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n            if torch.rand(1, device=device) < 0.45:\n                x = self._apply_feature_specific_noise(x)\n\n            # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n            if torch.rand(1, device=device) < 0.2:\n                x = self._apply_global_modulation(x)\n\n            # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n            if torch.rand(1, device=device) < 0.4:\n                x = self._apply_slow_drift(x)\n            # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n            if torch.rand(1, device=device) < 0.25:\n                x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n                    0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n            augmented_versions.append(x)\n\n        # Combine results\n        if num_augments == 1:\n            result = augmented_versions[0]\n        else:\n            # Stack along new dimension: (num_augments, B, T, C)\n            result = torch.stack(augmented_versions, dim=0)\n\n        # Restore original shape if needed\n        if num_augments == 1:\n            return self._restore_shape(result, squeeze)\n        else:\n            # For multiple augmentations, keep the extra dimension\n            if squeeze:\n                # If original was (T, C), we now have (num_augments, 1, T, C)\n                # We want to keep the num_augments dimension but remove the batch dim\n                return result.squeeze(1)\n            return result\n\n\nclass SyncTextNeuralAugment:\n    \"\"\"\n    Synchronized augmentation that randomly cuts segments from both neural signal and text.\n    Maintains alignment between neural features and text tokens.\n    \"\"\"\n\n    def __init__(self, min_cut_ratio: float = 0.1, max_cut_ratio: float = 0.4, p: float = 0.5):\n        \"\"\"\n        Args:\n            min_cut_ratio: Minimum ratio of sequence to cut out\n            max_cut_ratio: Maximum ratio of sequence to cut out  \n            p: Probability of applying this augmentation\n        \"\"\"\n        self.min_cut_ratio = min_cut_ratio\n        self.max_cut_ratio = max_cut_ratio\n        self.p = p\n\n    def __call__(self, neural: torch.Tensor, text: str) -> Tuple[torch.Tensor, str]:\n        \"\"\"\n        Apply synchronized random cutting to neural signal and text.\n\n        Args:\n            neural: Neural features tensor of shape (T, 512)\n            text: Original text string\n            tokenizer: Tokenizer to help align text with neural segments\n\n        Returns:\n            Cut neural tensor and corresponding text\n        \"\"\"\n        if random.random() > self.p or len(text.strip()) == 0:\n            return neural, text\n\n        T = neural.shape[0]\n        if T < 10:  # Too short to cut meaningfully\n            return neural, text\n\n        # Determine cut parameters\n        cut_ratio = random.uniform(self.min_cut_ratio, self.max_cut_ratio)\n        cut_length = max(1, int(T * cut_ratio))\n        start_idx = random.randint(0, T - cut_length)\n        end_idx = start_idx + cut_length\n\n        # Cut neural signal\n        cut_neural = torch.cat([neural[:start_idx], neural[end_idx:]], dim=0)\n\n        # Cut corresponding text - this is the tricky part\n        # We need to estimate text segments corresponding to neural time steps\n        # Simple approach: assume linear mapping (this may need refinement)\n        text_length = len(text)\n        if text_length > 0:\n            # Calculate text cut positions proportionally\n            text_start_ratio = start_idx / T\n            text_end_ratio = end_idx / T\n\n            text_start_idx = max(0, int(text_length * text_start_ratio))\n            text_end_idx = min(text_length, int(text_length * text_end_ratio))\n\n            # Cut text\n            cut_text = text[:text_start_idx] + text[text_end_idx:]\n        else:\n            cut_text = text\n\n        return cut_neural, cut_text\n\n\n# Updated dataset class with the new augmentation\nclass BrainToTextDataset(Dataset):\n    \"\"\"\n    Dataset for DSD-NLA\n    Load neural features + text labels (NO phonemes!)\n    \"\"\"\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        self.aug = PhysioAwareAugment()\n        self.sync_aug = SyncTextNeuralAugment(\n            max_cut_ratio=0.7, p=0.6)  # New synchronized augmentation 0.6 , 0.4\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 and text is not None:\n\n            # Apply existing physiological augmentation\n            neural = self.aug(neural)\n            # Apply synchronized augmentation to both neural and text\n            neural, text = self.sync_aug(neural, text)\n            # Update tokens after text modification\n            tokens = self.tokenizer.encode(text)\n\n        return {\n            \"neural\": neural,\n            \"tokens\": tokens,\n            \"text\": text,\n        }\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\"]\n                   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\n\nclass DSDNLALoss(nn.Module):\n    \"\"\"\n    Simplified loss:\n    - Only text generation loss (cross-entropy)\n    \"\"\"\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\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        ema_decay = 0.999  # 0.99\n        self.model_ema = ModelEmaV3(\n            self.model, decay=ema_decay, use_warmup=True, device=self.device\n        )\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        self.criterion = DSDNLALoss(ignore_index=0)\n        # self.criterion = HybridCTCAttentionLoss()\n        self.wer_loss = WordErrorRate()\n\n        # New AMP API\n\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.1,\n            anneal_strategy=\"cos\",\n        )\n\n        # self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        #     self.optimizer, patience=2, factor=0.98\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(\n                    self.device) if batch[\"tokens\"] is not None else None\n            )\n\n            # use_autocast = self.scaler is not None\n            # autocast_ctx = torch.amp.autocast(\"cuda\", enabled=use_autocast)\n            device = next(self.model.parameters()).device\n            device_type = device.type\n            self.optimizer.zero_grad(set_to_none=True)\n            with torch.autocast(device_type=device_type, dtype=torch.bfloat16):\n                logits = 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)\n                total_loss = losses[\"total\"]\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            # backward + clip + step (no scaler)\n            total_loss.backward()\n            torch.nn.utils.clip_grad_norm_(\n                self.model.parameters(), self.config.gradient_clip)\n            self.optimizer.step()\n\n            # ModelEmaV3 exposes .update(model)\n            self.model_ema.update(self.model)\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            self.global_step += 1\n\n        return float(np.mean(epoch_losses))\n\n    @torch.inference_mode()\n    def validate(self, val_loader: DataLoader) -> float:\n        self.model_ema.module.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(\n                    self.device) if batch[\"tokens\"] is not None else None\n            )\n\n            logits = self.model_ema.module(\n                neural,\n                None,\n                tokens,\n                neural_mask=neural_mask,\n                training=False,\n            )\n            losses = self.criterion(logits, tokens)\n            val_losses.append(losses[\"total\"].item())\n\n        return float(np.mean(val_losses))\n\n    @torch.inference_mode()\n    def wer(self, val_loader: DataLoader) -> float:\n        self.model.eval()\n        val_wer: List[float] = []\n\n        for batch in tqdm(val_loader, desc=\"Validation\"):\n            # 1. Get the full batch tensors\n            neural_batch = batch[\"neural\"].to(self.device)\n            mask_batch = batch[\"neural_mask\"].to(self.device)\n            texts_batch = batch[\"texts\"]  # List of strings\n\n            # 2. Iterate over every sample in the batch individually\n            for i in range(len(texts_batch)):\n                # Slice the tensor to get the i-th sample\n                # .unsqueeze(0) keeps dimensions as (1, Time, Channels)\n                # so the model still sees a \"batch\" of 1\n                single_neural = neural_batch[i].unsqueeze(0)\n                single_mask = mask_batch[i].unsqueeze(0)\n                single_text = texts_batch[i]\n\n                # 3. Generate for just this one sample\n                pred_text = self.model.generate_candidate(\n                    single_neural, single_mask)\n\n                # 4. Calculate WER for this pair\n                # Ensure wer_loss can handle (str, str). If it needs lists, use ([pred], [text])\n                wer_loss = self.wer_loss(pred_text, single_text)\n                val_wer.append(wer_loss)\n\n        return float(np.mean(val_wer))\n\n    def save_checkpoint(self, epoch: int):\n        # save ema as well\n        checkpoint = {\n            \"epoch\": epoch,\n            \"global_step\": self.global_step,\n            \"model_state_dict\": self.model.state_dict(),\n            \"ema_state\": self.model_ema.state_dict(),\n            \"optimizer_state_dict\": self.optimizer.state_dict(),\n            \"scheduler_state_dict\": (\n                self.scheduler.state_dict() if self.scheduler is not None else None\n            ),\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            # wer = self.wer(val_loader)\n            # print(f\"Epoch {epoch}: WER Loss = {wer:.4f}\")\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(val_loss)\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 >= 5:\n            #     self.model.freeze()\n            #     self.optimizer = optim.AdamW(\n            #         filter(lambda p: p.requires_grad, self.model.parameters()),\n            #         lr=config.learning_rate//2,\n            #         weight_decay=config.weight_decay,\n            #         betas=(0.9, 0.98),\n            #         eps=1e-8,\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:42:08.561723Z","iopub.execute_input":"2025-12-31T09:42:08.562063Z","iopub.status.idle":"2025-12-31T09:42:10.101160Z","shell.execute_reply.started":"2025-12-31T09:42:08.562031Z","shell.execute_reply":"2025-12-31T09:42:10.100557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gc.collect()\ntorch.cuda.empty_cache()","metadata":{"papermill":{"duration":1.649675,"end_time":"2025-12-21T14:46:23.477054","exception":false,"start_time":"2025-12-21T14:46:21.827379","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:48:08.742647Z","iopub.execute_input":"2025-12-31T09:48:08.743415Z","iopub.status.idle":"2025-12-31T09:48:10.332821Z","shell.execute_reply.started":"2025-12-31T09:48:08.743367Z","shell.execute_reply":"2025-12-31T09:48:10.332230Z"}},"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- Use Beam Search to generate top-k candidates\n- RESCORING: Use Qwen 3  with Adaptive Random-Word Detection\n- Normalize to strict WER format for submission\n\"\"\"\n\n# ============================================================================\n# 1. TEXT NORMALIZER — Strict WER Format (No punctuation)\n# ============================================================================\n\nimport os\nfrom transformers import pipeline\n\n# 1. Force the library into offline mode (prevents it from trying to check for updates)\nos.environ[\"TRANSFORMERS_OFFLINE\"] = \"1\"\n# ============================================================================\n# 2. MODERN POST-PROCESSING (Qwen 3 + Adaptive Logic)\n# ============================================================================\n\nprint(\"\\n\" + \"=\" * 50)\nprint(\"LOADING POST-PROCESSING TOOLS (Qwen 3)\")\nprint(\"=\" * 50)\n\nDEVICE_LM = \"cuda:0\" \nUSE_LM = True\n\ntry:\n    print(f\"Loading {config.LM_MODEL_ID}...\")\n    lm_tokenizer = AutoTokenizer.from_pretrained(config.LM_MODEL_ID)\n    lm_model = AutoModelForCausalLM.from_pretrained(\n        config.LM_MODEL_ID,\n        dtype=torch.bfloat16,  # FP16 is 2x faster and uses half RAM\n        device_map=DEVICE_LM,\n    )\n    lm_model.eval()\n    print(f\"✅ LM loaded successfully on {DEVICE_LM}\")\nexcept Exception as e:\n    print(f\"⚠️ Warning: Could not load Modern LM. Error: {e}\")\n    print(\"LM Rescoring will be DISABLED. Falling back to raw decoder.\")\n    lm_tokenizer = None\n    lm_model = None\n    USE_LM = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:42:10.457786Z","iopub.execute_input":"2025-12-31T09:42:10.458121Z","iopub.status.idle":"2025-12-31T09:42:13.879399Z","shell.execute_reply.started":"2025-12-31T09:42:10.458097Z","shell.execute_reply":"2025-12-31T09:42:13.878620Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# 2. Point to the LOCAL folder path\nLOCAL_PATH = \"/kaggle/input/sceb-b2t\"\n\n# 3. Load from the folder\n# 'local_files_only=True' ensures it won't crash trying to reach the internet\nspell_checker = pipeline(\n    \"text2text-generation\", \n    model=LOCAL_PATH, \n    tokenizer=LOCAL_PATH,\n)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:42:13.880451Z","iopub.execute_input":"2025-12-31T09:42:13.881038Z","iopub.status.idle":"2025-12-31T09:42:14.513415Z","shell.execute_reply.started":"2025-12-31T09:42:13.881008Z","shell.execute_reply":"2025-12-31T09:42:14.512817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import PretrainedConfig\nfrom transformers.modeling_outputs import Seq2SeqLMOutput\nfrom transformers.generation import GenerationMixin\nfrom transformers import PreTrainedModel, PretrainedConfig\nfrom typing import Optional\nfrom transformers import PretrainedConfig, PreTrainedModel\nfrom collections import defaultdict\nfrom typing import List, Tuple, Optional\nimport glob\n\nimport math\nfrom typing import List, Tuple, Optional, Dict, Any\nimport torch\nimport torch.nn.functional as F\nimport pandas as pd\nfrom pathlib import Path\n\nimport torch\nimport torch.nn.functional as F\nfrom typing import List, Tuple, Optional, Union\nimport pkg_resources\nfrom symspellpy import SymSpell, Verbosity\nfrom concurrent.futures import ThreadPoolExecutor\n\n\ndef calculate_sentence_score(text: str) -> float:\n    \"\"\"\n    Calculates the negative log-likelihood (NLL) per token.\n    Lower score = Better English sentence.\n    High score (> 6.0) = Likely 'Random Word' list or gibberish.\n    \"\"\"\n    if not USE_LM or not text.strip():\n        return 999.0\n\n    try:\n        inputs = lm_tokenizer(text, return_tensors=\"pt\").to(DEVICE_LM)\n        with torch.inference_mode():\n            outputs = lm_model(**inputs, labels=inputs[\"input_ids\"])\n            loss = outputs.loss.item()  # This is CrossEntropy (NLL)\n        return loss\n    except Exception as e:\n        return 999.0\n\n\ndef select_best_candidate_adaptive_batched(candidates: List[str]) -> str:\n    \"\"\"\n    Batched processing: Scores all 10 candidates in ONE GPU pass.\n    10x faster than looping.\n    \"\"\"\n    if not candidates:\n        return \"\"\n    if not USE_LM:\n        return candidates[0].strip()\n\n    # Filter empty strings\n    clean_candidates = [c.strip() for c in candidates if c.strip()]\n    if not clean_candidates:\n        return candidates[0]\n\n    try:\n        # Tokenize ALL candidates at once (Batch Size = Beam Width)\n        inputs = lm_tokenizer(\n            clean_candidates, return_tensors=\"pt\", padding=True, truncation=True\n        ).to(DEVICE_LM)\n\n        with torch.inference_mode():\n            # Forward pass returns loss for the whole batch\n            outputs = lm_model(\n                input_ids=inputs[\"input_ids\"],\n                attention_mask=inputs[\"attention_mask\"],\n                labels=inputs[\"input_ids\"],\n            )\n\n            # HuggingFace usually returns a scalar average loss for training.\n            # To get PER-SENTENCE loss, we need the raw logits.\n            # However, a faster hack for inference rescoring:\n            # We calculate loss manually to keep it batched.\n\n            # Shift logits and labels for causal LM loss\n            shift_logits = outputs.logits[..., :-1, :].contiguous()\n            shift_labels = inputs[\"input_ids\"][..., 1:].contiguous()\n            shift_mask = inputs[\"attention_mask\"][..., 1:].contiguous()\n\n            # Flatten to calculate cross entropy\n            PAD_TOKEN_ID = 0  # Replace with your actual pad token ID\n            loss_fct = torch.nn.CrossEntropyLoss(\n                reduction=\"none\", ignore_index=PAD_TOKEN_ID\n            )\n            token_losses = loss_fct(\n                shift_logits.view(-1, shift_logits.size(-1)\n                                  ), shift_labels.view(-1)\n            )\n\n            # Reshape back to (Batch, Seq_Len)\n            token_losses = token_losses.view(shift_labels.size())\n\n            # Apply attention mask so we don't count padding tokens in score\n            # Sum loss per sentence / Number of real tokens\n            sentence_sums = (token_losses * shift_mask).sum(dim=1)\n            sentence_lengths = shift_mask.sum(dim=1)\n\n            # Avoid division by zero\n            sentence_scores = sentence_sums / (sentence_lengths + 1e-9)\n\n            # Convert to list\n            scores = sentence_scores.tolist()\n\n    except Exception as e:\n        print(f\"Batch scoring failed: {e}. Falling back to beam 0.\")\n        return candidates[0].strip()\n\n    # Pair scores with text\n    scored_candidates = list(zip(scores, clean_candidates))\n\n    # Sort: Lowest Loss = Best\n    scored_candidates.sort(key=lambda x: x[0])\n\n    best_score, best_text = scored_candidates[0]\n    # Adaptive Threshold (Tuned for 3B/4B models)\n    # 3B models are confident. If loss > 5.0, it's definitely garbage/random.\n\n    return best_text  # Trust LM\n\n\n# ============================================================================\n# 4. INFERENCE ENGINE\n# ============================================================================\n\n\nclass InferenceEngine:\n    def __init__(self, model_path: str, device: str = \"cuda\", max_len: int = 200):\n        self.device = device\n        self.max_len = max_len\n        self.tokenizer = CharTokenizer()\n\n        ckpt_path = Path(model_path)\n        if not ckpt_path.exists():\n            raise FileNotFoundError(f\"Checkpoint not found at {ckpt_path}\")\n\n        print(f\"Loading Brain Model from {ckpt_path}...\")\n        checkpoint = torch.load(ckpt_path, map_location=device)\n        config = checkpoint.get(\"config\", {})\n\n        # IMPORTANT: Ensure DSDNLA class is available in scope or imported!\n        # Initializing simplified model based on config\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_encoder=config.get(\"n_heads_encoder\", 8),\n            n_heads_decoder=config.get(\"n_heads_decoder\", 8),\n            dropout_encoder=0.0,\n            dropout_decoder=0.0,\n        )\n        self.model.load_state_dict(checkpoint[\"model_state_dict\"])\n        self.model = self.model.to(device)\n        self.model.eval()\n\n        ema_decay = 0.999\n        self.model_ema = ModelEmaV3(\n            self.model, decay=ema_decay, use_warmup=True, device=device\n        )\n        self.model_ema.load_state_dict(checkpoint[\"ema_state\"])\n        self.model_ema.module.eval()\n\n    @torch.inference_mode()\n    def predict(self, neural_features) -> str:\n        \"\"\"Greedy prediction (fallback)\"\"\"\n        if not isinstance(neural_features, torch.Tensor):\n            neural_features = torch.tensor(\n                neural_features, dtype=torch.float32)\n        neural_features = neural_features.unsqueeze(0).to(self.device)\n        tokens = self.model.inference(neural_features, max_len=self.max_len)\n        return self.tokenizer.decode(tokens[0])\n\n\nclass MinimalDecoderConfig(PretrainedConfig):\n    is_encoder_decoder = True\n    is_decoder = True\n    num_hidden_layers = 1  # required by generation cache\n    vocab_size = 35  # just a placeholder, match your tokenizer\n    pad_token_id = 0\n    bos_token_id = 1\n    eos_token_id = 2\n\n\nclass DecoderHFWrapper(PreTrainedModel, GenerationMixin):\n    config_class = MinimalDecoderConfig\n\n    def __init__(self, decoder, tokenizer, config=None):\n        if config is None:\n            config = MinimalDecoderConfig()\n\n        config.is_encoder_decoder = True\n        config.vocab_size = tokenizer.vocab_size\n        config.bos_token_id = tokenizer.bos_id\n        config.eos_token_id = tokenizer.eos_id\n        config.pad_token_id = tokenizer.pad_id\n\n        super().__init__(config)\n\n        self.decoder = decoder\n        self.tokenizer = tokenizer\n\n    def forward(\n        self,\n        input_ids=None,\n        attention_mask=None,\n        encoder_outputs=None,\n        return_dict=True,\n        **kwargs,\n    ):\n        if isinstance(encoder_outputs, dict):\n            z_neural = encoder_outputs[\"last_hidden_state\"]\n        else:\n            z_neural = encoder_outputs\n        # apply ensemble here with corresponding z_neural\n        logits = self.decoder(z_neural, input_ids)\n\n        if return_dict:\n            return Seq2SeqLMOutput(logits=logits)\n        return (logits,)\n\n    def prepare_inputs_for_generation(\n        self,\n        input_ids,\n        encoder_outputs=None,\n        attention_mask=None,\n        **kwargs,\n    ):\n        return {\n            \"input_ids\": input_ids,\n            \"encoder_outputs\": encoder_outputs,\n            \"attention_mask\": attention_mask,\n        }\n\n\nclass EnsembleDecoderHFWrapper(PreTrainedModel, GenerationMixin):\n    config_class = MinimalDecoderConfig\n\n    def __init__(self, encoder_decoder_pairs, tokenizer, weights=None, config=None):\n        \"\"\"\n        Args:\n            encoder_decoder_pairs: List of tuples [(encoder1, decoder1), (encoder2, decoder2), ...]\n            tokenizer: Shared tokenizer\n            weights: Optional list of weights for each model (defaults to equal weights)\n        \"\"\"\n        if config is None:\n            config = MinimalDecoderConfig()\n\n        config.is_encoder_decoder = True\n        config.vocab_size = tokenizer.vocab_size\n        config.bos_token_id = tokenizer.bos_id\n        config.eos_token_id = tokenizer.eos_id\n        config.pad_token_id = tokenizer.pad_id\n\n        super().__init__(config)\n\n        self.encoders = nn.ModuleList([pair[0]\n                                      for pair in encoder_decoder_pairs])\n        self.decoders = nn.ModuleList([pair[1]\n                                      for pair in encoder_decoder_pairs])\n        self.tokenizer = tokenizer\n\n        # Default to equal weights if not provided\n        if weights is None:\n            self.weights = [1.0 / len(encoder_decoder_pairs)] * len(\n                encoder_decoder_pairs\n            )\n        else:\n            # Normalize weights to sum to 1\n            total = sum(weights)\n            self.weights = [w / total for w in weights]\n\n    def forward(\n        self,\n        input_ids=None,\n        attention_mask=None,\n        encoder_outputs=None,\n        return_dict=True,\n        **kwargs,\n    ):\n        \"\"\"\n        encoder_outputs should be a list of z_neural tensors, one for each model\n        or a single tensor that will be processed by all encoders\n        \"\"\"\n        # Handle different types of encoder_outputs input\n        if encoder_outputs is None:\n            raise ValueError(\n                \"encoder_outputs must be provided for ensemble models\")\n\n        # If encoder_outputs is a single tensor, process it through all encoders\n        if isinstance(encoder_outputs, torch.Tensor) or (\n            isinstance(encoder_outputs,\n                       dict) and \"last_hidden_state\" in encoder_outputs\n        ):\n            # This is a single input that needs to be processed by all encoders\n            if isinstance(encoder_outputs, dict):\n                neural_input = encoder_outputs[\"last_hidden_state\"]\n            else:\n                neural_input = encoder_outputs\n\n            # Get encoder outputs for each model\n            all_z_neural = []\n            for encoder in self.encoders:\n                # Assuming encoder returns a tensor or dict with \"last_hidden_state\"\n                z_neural, _ = encoder(neural_input)\n                all_z_neural.append(z_neural)\n            \n            # # Get encoder outputs for each model using threadpool\n            # def process_encoder(encoder):\n            #     z_neural, _ = encoder(neural_input)\n            #     return z_neural\n            \n            # with ThreadPoolExecutor(max_workers=4) as executor:\n            #     all_z_neural = list(executor.map(process_encoder, self.encoders))\n\n        \n        else:\n            raise ValueError(\n                \"encoder_outputs must be either a single tensor/dict or a list of \"\n                f\"{len(self.encoders)} tensors/dicts (one for each model)\"\n            )\n\n        # Get logits from each decoder using its corresponding encoder output\n        all_logits = []\n        for i, (decoder, z_neural) in enumerate(zip(self.decoders, all_z_neural)):\n            logits = decoder(z_neural, input_ids)\n            all_logits.append(logits)\n\n        \n        # # Process all decoders in parallel using threadpool\n        # def process_decoder_pair(decoder_z_pair):\n        #     decoder, z_neural = decoder_z_pair\n        #     return decoder(z_neural, input_ids)\n        \n        # decoder_z_pairs = list(zip(self.decoders, all_z_neural))\n        # with ThreadPoolExecutor(max_workers=4) as executor:\n        #     all_logits = list(executor.map(process_decoder_pair, decoder_z_pairs))\n\n\n        \n        # # Weighted average of logits\n        weighted_logits = None\n        for i, logits in enumerate(all_logits):\n            if weighted_logits is None:\n                weighted_logits = self.weights[i] * logits\n            else:\n                weighted_logits += self.weights[i] * logits\n\n        # Convert logits to probabilities using softmax (with numerical stability)\n        # all_probs = []\n        # for logits in all_logits:\n        #     # Apply numerical stability trick: subtract max value before exponentiation\n        #     max_logits = logits.max(dim=-1, keepdim=True)[0]\n        #     stabilized_logits = logits - max_logits\n        #     probs = torch.exp(stabilized_logits) / \\\n        #         torch.exp(stabilized_logits).sum(dim=-1, keepdim=True)\n        #     all_probs.append(probs)\n\n        # # Weighted average of probabilities\n        # weighted_probs = None\n        # for i, probs in enumerate(all_probs):\n        #     if weighted_probs is None:\n        #         weighted_probs = self.weights[i] * probs\n        #     else:\n        #         weighted_probs += self.weights[i] * probs\n\n        # # Normalize to ensure it's a proper probability distribution\n        # weighted_probs = weighted_probs / \\\n        #     weighted_probs.sum(dim=-1, keepdim=True)\n\n        # # Convert back to logits using log (with small epsilon for numerical stability)\n        # epsilon = 1e-12  # Prevent log(0) which would be -infinity\n        # weighted_logits = torch.log(weighted_probs + epsilon)\n\n        if return_dict:\n            return Seq2SeqLMOutput(logits=weighted_logits)\n        return (weighted_logits,)\n\n    def prepare_inputs_for_generation(\n        self,\n        input_ids,\n        encoder_outputs=None,\n        attention_mask=None,\n        **kwargs,\n    ):\n        return {\n            \"input_ids\": input_ids,\n            \"encoder_outputs\": encoder_outputs,\n            \"attention_mask\": attention_mask,\n        }\n\n\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    beam_width: int = 15,\n) -> pd.DataFrame:\n\n    # Setup\n    nen = 10\n    engine = [0] * nen\n    for fold in range(nen):\n        engine[fold] = InferenceEngine(\n            os.path.join(\n                model_path, f\"best_dsdnla_model_{fold}.pt\"),\n            device=device,\n        )\n\n    encoder_decoder_pairs = [\n        (engine[i].model_ema.module.neural_encoder,\n         engine[i].model_ema.module.text_decoder)\n        for i in range(nen)\n    ]\n\n    wrapper = EnsembleDecoderHFWrapper(\n        encoder_decoder_pairs, engine[0].tokenizer\n    ).to(device)\n\n    # aug = PhysioAwareAugment()\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    # Create a SymSpell object\n    sym_spell = SymSpell(max_dictionary_edit_distance=2, prefix_length=9)\n\n    # Load the dictionary\n    # This path points to the dictionary file installed with the package\n    dictionary_path = pkg_resources.resource_filename(\n        \"symspellpy\", \"frequency_dictionary_en_82_765.txt\"\n    )\n    sym_spell.load_dictionary(dictionary_path, term_index=0, count_index=1)\n\n    def correct_sentence_word_by_word(sentence):\n        \n        results = spell_checker(sentence, max_new_tokens=100, num_beams=5, early_stopping=True)\n\n        return results[0]['generated_text']\n\n    print(\"\\nStarting Inference with Adaptive Rescoring...\")\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            # Iterate trials\n            for key in tqdm(keys, leave=False):\n                neural_features = f[key][\"input_features\"][:]  # (T, 256)\n                # 1. Encode Neural Data\n                neural_tensor = torch.tensor(\n                    neural_features).unsqueeze(0).to(device)\n                # z_neural, _ = engine.model_ema.module.neural_encoder(neural_tensor)\n                # z_neural_aug = aug.forward_tta(z_neural, num_augments=10)\n                # 2. Generate Candidates (Beam Search)\n                # candidates = beam_search.generate_candidates(z_neural)\n                # # TTA to generate more candidates could be added here\n                # print(candidates)\n                # Get all candidate texts\n                all_candidates = []\n                # for z_n in z_neural_aug:\n                outputs = wrapper.generate(\n                    input_ids=torch.tensor(\n                        [[engine[0].tokenizer.bos_id]], device=device),\n                    encoder_outputs={\"last_hidden_state\": neural_tensor},\n                    num_beams=7,\n                    num_return_sequences=5,\n                    max_new_tokens=100,\n                    repetition_penalty=1.0,\n                    early_stopping=True,\n                    output_scores=True,\n                    return_dict_in_generate=True,\n                )\n\n                # (num_return_sequences, seq_len)\n                sequences = outputs.sequences\n                # (num_return_sequences,)\n                scores = outputs.sequences_scores\n\n                # Convert each sequence to a tuple (score, token_ids as list)\n                for seq, score in zip(sequences, scores):\n                    all_candidates.append((score.item(), seq.tolist()))\n\n                # Sort by score descending\n                # all_candidates.sort(key=lambda x: x[0], reverse=True)\n\n                # # Take top-k\n                # top_k = 5\n                # top_candidates = all_candidates[:top_k]\n\n                # Decode\n                candidates = [engine[0].tokenizer.decode(seq)\n                              for score, seq in all_candidates]\n                \n                candidates = list(set(candidates))\n\n                # candidates = [correct_sentence_word_by_word(s) for s in candidates]\n                # 3. Adaptive Selection (The \"Secret Sauce\")\n                # Decides whether to use Qwen or Raw Decoder based on likelihood\n                best_text = select_best_candidate_adaptive_batched(candidates)\n                best_text = correct_sentence_word_by_word(best_text)\n                best_text = normalize_for_eval(clean_generated_text(best_text))\n                print(best_text)\n\n                all_predictions.append(best_text)\n\n    # Save\n    submission_df = pd.DataFrame(\n        {\"id\": range(len(all_predictions)), \"text\": all_predictions}\n    )\n    submission_df.to_csv(output_path, index=False)\n\n    print(f\"\\n✅ Submission saved to {output_path}\")\n    print(\"Example Output:\")\n    print(submission_df.head(10))\n    return submission_df\n\n\n# ============================================================================\n# RUNNER\n# ============================================================================\n\nif __name__ == \"__main__\":\n\n    # Ensure model exists before running\n    if Path(config.MODEL_PATH).exists():\n        generate_submission(\n            model_path=config.MODEL_PATH,\n            test_data_dir=config.TEST_DIR,\n            output_path=config.OUTPUT_CSV,\n            beam_width=10,  # Higher beam width = better chance for LM to find good sentences\n            device=f\"cuda:{config.gpu_id}\"\n\n        )\n    else:\n        print(f\"❌ Model not found at {config.MODEL_PATH}. Run training first.\")","metadata":{"papermill":{"duration":3055.317032,"end_time":"2025-12-21T15:37:20.108557","exception":false,"start_time":"2025-12-21T14:46:24.791525","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:48:17.681780Z","iopub.execute_input":"2025-12-31T09:48:17.682464Z","iopub.status.idle":"2025-12-31T09:49:03.957704Z","shell.execute_reply.started":"2025-12-31T09:48:17.682432Z","shell.execute_reply":"2025-12-31T09:49:03.956582Z"}},"outputs":[],"execution_count":null}]}