{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"**PHASE 1 OF THE PROJECT**","metadata":{}},{"cell_type":"code","source":"# Cell A: install (h5py) and imports, device check\n!pip install --quiet h5py einops jiwer\n\nimport os, sys, math, json, time, glob\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport h5py\n\nimport torch\nprint(\"Torch\", torch.__version__, \"CUDA available:\", torch.cuda.is_available())\nprint(\"Device visible:\", os.environ.get(\"COLAB_TPU_ADDR\", \"no TPU env var; running on GPU/CPU\"))\n# In Kaggle you can choose accelerator in the notebook settings. Pick \"GPU\" -> P100 or T4.\n# show selected GPU name if available\nif torch.cuda.is_available():\n    try:\n        import subprocess\n        gpu_info = subprocess.check_output([\"nvidia-smi\", \"--query-gpu=name,memory.total\", \"--format=csv,noheader\"]).decode().strip()\n        print(\"GPU info:\", gpu_info)\n    except Exception as e:\n        print(\"Could not run nvidia-smi:\", e)\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:14:28.145210Z","iopub.execute_input":"2025-11-07T05:14:28.145430Z","iopub.status.idle":"2025-11-07T05:14:43.015483Z","shell.execute_reply.started":"2025-11-07T05:14:28.145407Z","shell.execute_reply":"2025-11-07T05:14:43.014596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell B: list the dataset tree (show relevant hdf5 files)\nINPUT_DIR = Path('/kaggle/input/brain-to-text-25')\nprint(\"Input root exists?\", INPUT_DIR.exists())\n# get all hdf5 files recursively under the input dir\nhdf5_files = sorted([p for p in INPUT_DIR.rglob(\"*.hdf5\")])\nprint(f\"Found {len(hdf5_files)} .hdf5 files (showing first 40):\")\nfor p in hdf5_files[:40]:\n    print(\"-\", p.relative_to(INPUT_DIR))\n# If the dataset is in a subfolder (like t15_copyTask_neuralData/hdf5_data_final/...), print that too\ntop_subdirs = [d for d in INPUT_DIR.iterdir() if d.is_dir()]\nprint(\"\\nTop-level folders in input:\")\nfor d in top_subdirs:\n    print(\"-\", d.name)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:14:43.016385Z","iopub.execute_input":"2025-11-07T05:14:43.016816Z","iopub.status.idle":"2025-11-07T05:14:43.465868Z","shell.execute_reply.started":"2025-11-07T05:14:43.016794Z","shell.execute_reply":"2025-11-07T05:14:43.465190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell C: open one train hdf5 and list its dataset keys and shapes\nif len(hdf5_files) == 0:\n    raise SystemExit(\"No .hdf5 files found under /kaggle/input/brain-to-text-25. Please check dataset path in the notebook UI.\")\n\n# choose a representative train file if possible (prefer names containing 'train' or 'data_train')\ntrain_candidates = [p for p in hdf5_files if 'train' in p.name.lower() or 'data_train' in p.name.lower()]\nsample_file = train_candidates[0] if train_candidates else hdf5_files[0]\nprint(\"Using sample file for inspection:\", sample_file.relative_to(INPUT_DIR))\n\nwith h5py.File(sample_file, 'r') as hf:\n    print(\"Top-level keys:\", list(hf.keys()))\n    # iterate and print shapes/types for datasets and groups\n    def print_group(g, indent=0):\n        for k in g:\n            item = g[k]\n            if isinstance(item, h5py.Dataset):\n                print(\"  \" * indent + f\"- Dataset: {k}  shape={item.shape}  dtype={item.dtype}\")\n            elif isinstance(item, h5py.Group):\n                print(\"  \" * indent + f\"- Group: {k}\")\n                print_group(item, indent+1)\n    print_group(hf)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:14:43.467326Z","iopub.execute_input":"2025-11-07T05:14:43.467543Z","iopub.status.idle":"2025-11-07T05:14:45.726218Z","shell.execute_reply.started":"2025-11-07T05:14:43.467525Z","shell.execute_reply":"2025-11-07T05:14:45.725539Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell D: read one example, print transcript and plot a few channels/time-range\nwith h5py.File(sample_file, 'r') as hf:\n    # Try to find likely dataset names that hold signals and labels\n    # Common keys might be: 'neural_activity', 'signals', 'data', 'transcripts', 'labels', 'text'\n    keys = list(hf.keys())\n    print(\"File keys:\", keys)\n\n    # Heuristic: find first dataset with 2D shape (time x channels) or 3D (samples x time x channels)\n    ds_candidates = []\n    for k in keys:\n        item = hf[k]\n        if isinstance(item, h5py.Dataset):\n            ds_candidates.append((k, item.shape))\n        elif isinstance(item, h5py.Group):\n            # look inside group\n            for kk in item.keys():\n                it = item[kk]\n                if isinstance(it, h5py.Dataset):\n                    ds_candidates.append((f\"{k}/{kk}\", it.shape))\n    print(\"Dataset candidates (name, shape):\")\n    for name, shape in ds_candidates[:40]:\n        print(\"-\", name, shape)\n\n    # Now try common names\n    def try_get(path):\n        try:\n            return hf[path]\n        except Exception:\n            return None\n\n    # Attempt patterns\n    signal_ds = None\n    transcript_ds = None\n    # search dataset names for 'signal' or 'neural' or 'data'\n    for name, shape in ds_candidates:\n        lname = name.lower()\n        if 'neural' in lname or 'signal' in lname or 'data' in lname or 'activity' in lname:\n            # choose a candidate dataset with at least 2 dims\n            arr = hf[name]\n            if len(arr.shape) >= 2:\n                signal_ds = name\n                break\n    # search for transcript/text label datasets\n    for name, shape in ds_candidates:\n        lname = name.lower()\n        if 'trans' in lname or 'label' in lname or 'text' in lname or 'target' in lname:\n            transcript_ds = name\n            break\n\n    print(\"Guessed signal dataset:\", signal_ds)\n    print(\"Guessed transcript dataset:\", transcript_ds)\n\n    # If signal_ds is a 3D dataset (N x T x C), load first sample; if 2D, treat as T x C\n    if signal_ds is None:\n        # fallback: pick the first dataset with 2+ dims\n        for name, shape in ds_candidates:\n            if len(shape) >= 2:\n                signal_ds = name\n                break\n    if signal_ds is None:\n        raise SystemExit(\"Couldn't locate a signal dataset automatically. Please inspect the file keys printed above.\")\n\n    sig = hf[signal_ds]\n    print(\"signal dataset shape:\", sig.shape)\n    if len(sig.shape) == 3:\n        # (N, T, C) or (N, C, T) - check typical ordering\n        sample_signal = sig[0]  # first sample\n    elif len(sig.shape) == 2:\n        sample_signal = sig[:]   # whole array (T x C) or (samples x something) — we'll infer\n    else:\n        raise SystemExit(\"Signal dataset has unexpected number of dims:\", sig.shape)\n\n    # convert to numpy\n    sample_signal = np.array(sample_signal)\n    print(\"Loaded sample_signal shape:\", sample_signal.shape)\n\n    # If sample_signal shape is (T, C) we are good. If it's (C, T) transpose.\n    if sample_signal.shape[0] < sample_signal.shape[1] and sample_signal.shape[0] < 10:\n        # likely (channels, time) -> transpose\n        sample_signal = sample_signal.T\n        print(\"Transposed sample to shape (time, channels):\", sample_signal.shape)\n\n    # Print min/max and a small slice\n    print(\"sample_signal dtype:\", sample_signal.dtype, \"min/max:\", sample_signal.min(), sample_signal.max())\n    t_len, n_ch = sample_signal.shape\n    print(f\"Time steps: {t_len}, Channels: {n_ch}\")\n\n    # Plot a few channels (first 3 channels) across a window (first 2000 timesteps or full if shorter)\n    plt.figure(figsize=(12,4))\n    ts = min(2000, t_len)\n    for ch in range(min(3, n_ch)):\n        plt.plot(sample_signal[:ts, ch] + ch* (np.std(sample_signal[:,ch])*4), label=f\"ch{ch}\")\n    plt.title(f\"Sample signal (first {ts} timesteps) — channels offset for visibility\")\n    plt.legend()\n    plt.xlabel(\"time\")\n    plt.show()\n\n    # Show transcript/target if found\n    if transcript_ds:\n        tr = np.array(hf[transcript_ds])\n        print(\"Transcript dataset shape:\", tr.shape)\n        # try to print a first transcript if dataset is list-like\n        try:\n            # h5py can store variable-length strings; decode if bytes\n            first = tr[0]\n            if isinstance(first, bytes):\n                first = first.decode('utf-8', errors='replace')\n            print(\"First transcript (raw):\", first)\n        except Exception as e:\n            print(\"Could not read a transcript entry directly:\", e)\n    else:\n        print(\"No transcript dataset automatically found in this file (we'll search other files).\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:14:45.726972Z","iopub.execute_input":"2025-11-07T05:14:45.727226Z","iopub.status.idle":"2025-11-07T05:14:47.927343Z","shell.execute_reply.started":"2025-11-07T05:14:45.727207Z","shell.execute_reply":"2025-11-07T05:14:47.926543Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**PHASE 2 OF THE PROJECT**","metadata":{}},{"cell_type":"code","source":"# Cell 1: Build an index mapping each trial -> (hdf5_path, group_name)\nfrom pathlib import Path\nimport h5py\nimport pandas as pd\nimport json, tqdm\n\nINPUT_DIR = Path('/kaggle/input/brain-to-text-25') / 't15_copyTask_neuralData' / 'hdf5_data_final'\nassert INPUT_DIR.exists(), f\"Expected dataset at {INPUT_DIR}, but not found.\"\n\n# find all train files (files named like data_train.hdf5)\nall_h5 = sorted([p for p in INPUT_DIR.rglob(\"data_train.hdf5\")])\nprint(\"Found train hdf5 files:\", len(all_h5))\nall_h5[:5]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:14:47.928557Z","iopub.execute_input":"2025-11-07T05:14:47.929198Z","iopub.status.idle":"2025-11-07T05:14:48.046843Z","shell.execute_reply.started":"2025-11-07T05:14:47.929177Z","shell.execute_reply":"2025-11-07T05:14:48.046269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 2: Create a DataFrame index where each row = one trial\nrows = []\nfor h5path in tqdm.tqdm(all_h5, desc=\"Scanning train hdf5 files\"):\n    with h5py.File(h5path, 'r') as hf:\n        # list group names that start with 'trial_'\n        for grp_name in hf.keys():\n            if grp_name.startswith('trial_'):\n                # optional: verify it has input_features and transcription/seq_class_ids\n                grp = hf[grp_name]\n                if 'input_features' in grp:\n                    # we won't load data now, just record shapes if available\n                    try:\n                        tshape = grp['input_features'].shape\n                    except Exception:\n                        tshape = None\n                    rows.append({\n                        'h5_path': str(h5path),\n                        'group': grp_name,\n                        'feat_shape': tshape\n                    })\n\nidx_df = pd.DataFrame(rows)\nprint(\"Total trials indexed:\", len(idx_df))\ndisplay(idx_df.head())\n# Save index for debugging / reproducibility\nidx_df.to_csv('/kaggle/working/trial_index.csv', index=False)\nprint(\"Saved /kaggle/working/trial_index.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:14:48.047589Z","iopub.execute_input":"2025-11-07T05:14:48.047921Z","iopub.status.idle":"2025-11-07T05:15:19.449163Z","shell.execute_reply.started":"2025-11-07T05:14:48.047867Z","shell.execute_reply":"2025-11-07T05:15:19.448591Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 4: collate_fn that returns: x [B,C,T], targets_concat [sum(L)], input_lengths [B], target_lengths [B]\nimport torch\nimport torch.nn as nn\n\ndef collate_for_ctc(batch):\n    \"\"\"\n    batch: list of tuples (feats_tensor [T,C], target_tensor [L])\n    Returns:\n      x_padded: [B, C, T_max]\n      targets_concat: 1D tensor of concatenated targets\n      input_lengths: tensor [B] lengths (in frames, i.e., T_i)\n      target_lengths: tensor [B] lengths (L_i)\n    \"\"\"\n    xs, ys = zip(*batch)\n    x_lens = [x.shape[0] for x in xs]\n    t_lens = [y.shape[0] for y in ys]\n\n    # pad inputs along time to max_t\n    max_t = max(x_lens)\n    channels = xs[0].shape[1]\n    x_padded = torch.zeros(len(xs), channels, max_t, dtype=torch.float32)\n    for i, x in enumerate(xs):\n        T = x.shape[0]\n        # x is [T, C] -> convert to [C, T]\n        x_padded[i, :, :T] = x.permute(1,0)\n\n    # concatenate targets to 1D for CTC\n    if sum(t_lens) > 0:\n        targets_concat = torch.cat([y.to(torch.long) for y in ys])\n    else:\n        targets_concat = torch.tensor([], dtype=torch.long)\n\n    return x_padded, targets_concat, torch.tensor(x_lens, dtype=torch.long), torch.tensor(t_lens, dtype=torch.long)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:15:19.449828Z","iopub.execute_input":"2025-11-07T05:15:19.450045Z","iopub.status.idle":"2025-11-07T05:15:19.456297Z","shell.execute_reply.started":"2025-11-07T05:15:19.450029Z","shell.execute_reply":"2025-11-07T05:15:19.455682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Re-define BrainDataset class (safe standalone cell)\nimport torch, h5py, collections, numpy as np\nfrom torch.utils.data import Dataset\n\nclass BrainDataset(Dataset):\n    def __init__(self, index_df, cache_size=8, max_len=None):\n        self.df = index_df.reset_index(drop=True)\n        self.max_len = max_len\n        self._cache_size = cache_size\n        self._file_cache = collections.OrderedDict()\n\n    def __len__(self):\n        return len(self.df)\n\n    def _open_file(self, path):\n        if path in self._file_cache:\n            self._file_cache.move_to_end(path)\n            return self._file_cache[path]\n        f = h5py.File(path, 'r')\n        self._file_cache[path] = f\n        if len(self._file_cache) > self._cache_size:\n            old_path, old_f = self._file_cache.popitem(last=False)\n            try: old_f.close()\n            except: pass\n        return f\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        f = self._open_file(row['h5_path'])\n        g = f[row['group']]\n\n        feats = g['input_features'][()].astype('float32')\n        if 'transcription' in g:\n            tgt = g['transcription'][()]\n        elif 'seq_class_ids' in g:\n            tgt = g['seq_class_ids'][()]\n        else:\n            raise KeyError(f\"No target found in {row['h5_path']}::{row['group']}\")\n\n        tgt = np.array(tgt, dtype='int64').reshape(-1)\n        if tgt.shape[0] >= 64:\n            nz = np.nonzero(tgt)[0]\n            tgt = tgt[:nz[-1]+1] if nz.size else tgt[:1]\n\n        if self.max_len and feats.shape[0] > self.max_len:\n            start = (feats.shape[0] - self.max_len)//2\n            feats = feats[start:start+self.max_len]\n\n        return torch.from_numpy(feats), torch.from_numpy(tgt)\n\n    def close(self):\n        for _, f in list(self._file_cache.items()):\n            try: f.close()\n            except: pass\n        self._file_cache.clear()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:15:19.457115Z","iopub.execute_input":"2025-11-07T05:15:19.457477Z","iopub.status.idle":"2025-11-07T05:15:19.484353Z","shell.execute_reply.started":"2025-11-07T05:15:19.457445Z","shell.execute_reply":"2025-11-07T05:15:19.483787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 5: split index into train/val by selecting files (we indexed only data_train files earlier)\n# But there are also separate data_val.hdf5 files; for validation we'll create a val index similarly by scanning data_val files.\n\n# Build val index (scan data_val.hdf5)\nval_h5 = sorted([p for p in INPUT_DIR.rglob(\"data_val.hdf5\")])\nprint(\"Found val hdf5 files:\", len(val_h5))\n\nval_rows = []\nfor h5path in val_h5:\n    with h5py.File(h5path, 'r') as hf:\n        for grp_name in hf.keys():\n            if grp_name.startswith('trial_') and 'input_features' in hf[grp_name]:\n                val_rows.append({'h5_path': str(h5path), 'group': grp_name})\nval_df = pd.DataFrame(val_rows)\nprint(\"Total val trials indexed:\", len(val_df))\n\n# Use the train index we built earlier (idx_df saved previously)\ntrain_df_index = idx_df  # from earlier cell\nprint(\"Train trials:\", len(train_df_index))\n\n# Create dataset objects (optionally set max_len to a value to limit memory; we can keep None to use full length)\ntrain_ds = BrainDataset(train_df_index, cache_size=12, max_len=None)\nval_ds = BrainDataset(val_df, cache_size=4, max_len=None)\n\nfrom torch.utils.data import DataLoader\ntrain_loader = DataLoader(train_ds, batch_size=8, shuffle=True, collate_fn=collate_for_ctc, num_workers=2)\nval_loader = DataLoader(val_ds, batch_size=8, shuffle=False, collate_fn=collate_for_ctc, num_workers=2)\n\n# sanity-check: load one batch\nbatch = next(iter(train_loader))\nx_padded, targets_concat, input_lengths, target_lengths = batch\nprint(\"x_padded shape (B,C,T):\", x_padded.shape)\nprint(\"targets_concat shape (sumL,):\", targets_concat.shape)\nprint(\"input_lengths:\", input_lengths)\nprint(\"target_lengths:\", target_lengths)\n\n# close datasets' file caches (they remain usable after reopening)\ntrain_ds.close()\nval_ds.close()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:15:19.486491Z","iopub.execute_input":"2025-11-07T05:15:19.486677Z","iopub.status.idle":"2025-11-07T05:15:26.431200Z","shell.execute_reply.started":"2025-11-07T05:15:19.486663Z","shell.execute_reply":"2025-11-07T05:15:26.430290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**PHASE 3 OF THE PROJECT**","metadata":{}},{"cell_type":"code","source":"# Phase 3 - Cell 1: Define the Transformer model architecture\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport math\n\n# --- Positional Encoding (standard Transformer-style sine/cosine) ---\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_len=2000):\n        super().__init__()\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe)\n\n    def forward(self, x):\n        # x shape: [B, T, D]\n        x = x + self.pe[:x.size(1)].unsqueeze(0)\n        return x\n\n\n# --- Convolutional Frontend (feature extractor + downsampler) ---\nclass ConvStem(nn.Module):\n    def __init__(self, in_ch, d_model):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv1d(in_ch, d_model // 2, kernel_size=7, stride=2, padding=3),\n            nn.ReLU(),\n            nn.Conv1d(d_model // 2, d_model, kernel_size=5, stride=2, padding=2),\n            nn.ReLU(),\n        )\n\n    def forward(self, x):\n        # input: [B, C, T]\n        return self.net(x)  # [B, D, T']\n\n\n# --- Full Brain-to-Text Model ---\nclass BrainToTextModel(nn.Module):\n    def __init__(self, in_ch=512, d_model=384, nhead=8, num_layers=6, vocab_size=200):\n        super().__init__()\n        self.conv = ConvStem(in_ch, d_model)\n        encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead,\n                                                   dim_feedforward=d_model * 4, dropout=0.1, activation='relu')\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n        self.pos_enc = PositionalEncoding(d_model)\n        self.fc = nn.Linear(d_model, vocab_size)\n        self.apply(self._init_weights)\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Linear):\n            nn.init.xavier_uniform_(m.weight)\n            if m.bias is not None:\n                nn.init.zeros_(m.bias)\n        elif isinstance(m, nn.Conv1d):\n            nn.init.kaiming_uniform_(m.weight, nonlinearity='relu')\n\n    def forward(self, x):\n        # x: [B, C, T]\n        x = self.conv(x)          # [B, D, T']\n        x = x.permute(0, 2, 1)    # [B, T', D]\n        x = self.pos_enc(x)\n        x = x.permute(1, 0, 2)    # [T', B, D] (required by Transformer)\n        x = self.transformer(x)   # [T', B, D]\n        x = x.permute(1, 0, 2)    # [B, T', D]\n        logits = self.fc(x)       # [B, T', vocab_size]\n        log_probs = F.log_softmax(logits, dim=-1)\n        return log_probs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:15:26.432289Z","iopub.execute_input":"2025-11-07T05:15:26.432534Z","iopub.status.idle":"2025-11-07T05:15:26.443771Z","shell.execute_reply.started":"2025-11-07T05:15:26.432513Z","shell.execute_reply":"2025-11-07T05:15:26.443219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 3 - Cell 2: Initialize model, optimizer, scheduler, loss\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.cuda.amp import GradScaler\n\n# detect input channels from one batch\nsample_batch = next(iter(train_loader))\nx_padded, targets_concat, input_lengths, target_lengths = sample_batch\nin_channels = x_padded.shape[1]\nprint(\"Detected input channels:\", in_channels)\n\n# define model hyperparameters\nvocab_size_estimate = 256  # since labels are already integer-encoded\nmodel = BrainToTextModel(in_ch=in_channels, d_model=384, nhead=8,\n                         num_layers=6, vocab_size=vocab_size_estimate).to('cuda')\n\n# define loss, optimizer, scheduler\nctc_loss = nn.CTCLoss(blank=0, zero_infinity=True)\noptimizer = AdamW(model.parameters(), lr=3e-4)\nscheduler = CosineAnnealingLR(optimizer, T_max=10)\nscaler = GradScaler()\n\nprint(\"Model initialized successfully.\")\nprint(\"Total parameters:\", sum(p.numel() for p in model.parameters())/1e6, \"Million\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:15:26.444591Z","iopub.execute_input":"2025-11-07T05:15:26.444820Z","iopub.status.idle":"2025-11-07T05:15:31.134671Z","shell.execute_reply.started":"2025-11-07T05:15:26.444792Z","shell.execute_reply":"2025-11-07T05:15:31.133829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**PHASE 4 OF THE PROJECT**","metadata":{}},{"cell_type":"code","source":"# =========================================================\n# Phase 4 - Cell 1 (FINAL VERSION)\n# Training & Validation loop with correct autocast + input_lengths fix\n# =========================================================\n\nimport time, torch\nfrom torch.amp import autocast\nfrom tqdm import tqdm\nimport jiwer\n\n\n# ------------------------ TRAIN ONE EPOCH ------------------------\ndef train_one_epoch(model, loader, optimizer, scheduler, scaler, epoch, device='cuda'):\n    model.train()\n    total_loss, steps = 0.0, 0\n    start_time = time.time()\n\n    for batch in tqdm(loader, desc=f\"Epoch {epoch} [train]\", leave=False):\n        x, targets_concat, input_lengths, target_lengths = batch\n        x, targets_concat = x.to(device), targets_concat.to(device)\n        input_lengths, target_lengths = input_lengths.to(device), target_lengths.to(device)\n\n        # 🔧 Fix for ConvStem downsampling (stride 2 twice → /4)\n        input_lengths = torch.div(input_lengths, 4, rounding_mode='floor')\n\n        optimizer.zero_grad(set_to_none=True)\n        with autocast('cuda', dtype=torch.float16):\n            log_probs = model(x)                     # [B, T', V]\n            log_probs = log_probs.permute(1, 0, 2)   # [T', B, V] for CTC\n            loss = ctc_loss(log_probs, targets_concat, input_lengths, target_lengths)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n\n        total_loss += loss.item()\n        steps += 1\n\n    avg_loss = total_loss / max(1, steps)\n    print(f\"Epoch {epoch}: train loss = {avg_loss:.4f}  (time {time.time() - start_time:.1f}s)\")\n    return avg_loss\n\n\n# ------------------------ VALIDATE ------------------------\n@torch.no_grad()\ndef validate(model, loader, epoch, device='cuda'):\n    model.eval()\n    total_loss, steps = 0.0, 0\n    preds, refs = [], []\n\n    for batch in tqdm(loader, desc=f\"Epoch {epoch} [val]\", leave=False):\n        x, targets_concat, input_lengths, target_lengths = batch\n        x, targets_concat = x.to(device), targets_concat.to(device)\n        input_lengths, target_lengths = input_lengths.to(device), target_lengths.to(device)\n\n        # 🔧 Adjust lengths for ConvStem downsampling\n        input_lengths = torch.div(input_lengths, 4, rounding_mode='floor')\n\n        log_probs = model(x)\n        log_probs = log_probs.permute(1, 0, 2)\n        loss = ctc_loss(log_probs, targets_concat, input_lengths, target_lengths)\n        total_loss += loss.item()\n        steps += 1\n\n        # --- Greedy decode for approximate WER ---\n        decoded_batch = log_probs.argmax(-1).permute(1, 0).cpu().numpy()\n        idx = 0\n        for i, L in enumerate(target_lengths):\n            true_seq = targets_concat[idx:idx + L].cpu().numpy().tolist()\n            idx += L\n            pred_seq = decoded_batch[i]\n            # collapse repeats + remove blanks\n            pred_seq_clean = [p for j, p in enumerate(pred_seq) if (j == 0 or p != pred_seq[j - 1]) and p != 0]\n            ref_seq_clean = [p for p in true_seq if p != 0]\n            preds.append(\" \".join(map(str, pred_seq_clean)))\n            refs.append(\" \".join(map(str, ref_seq_clean)))\n\n    val_loss = total_loss / max(1, steps)\n    try:\n        wer_score = jiwer.wer(refs, preds)\n    except Exception:\n        wer_score = None\n\n    print(f\"Epoch {epoch}: val loss = {val_loss:.4f},  WER ≈ {wer_score if wer_score else 'N/A'}\")\n    return val_loss, wer_score\n\n\n# ------------------------ SAVE CHECKPOINT ------------------------\ndef save_checkpoint(model, optimizer, scheduler, epoch, path):\n    ckpt = {\n        'epoch': epoch,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'scheduler_state_dict': scheduler.state_dict(),\n    }\n    torch.save(ckpt, path)\n    print(f\"✅ Saved checkpoint: {path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:15:31.135717Z","iopub.execute_input":"2025-11-07T05:15:31.136535Z","iopub.status.idle":"2025-11-07T05:15:31.170321Z","shell.execute_reply.started":"2025-11-07T05:15:31.136507Z","shell.execute_reply":"2025-11-07T05:15:31.169524Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Phase 4 - Cell 2: Main training driver\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nnum_epochs = 5  # test run first\nbest_val_loss = float('inf')\n\nfor epoch in range(1, num_epochs + 1):\n    train_loss = train_one_epoch(model, train_loader, optimizer, scheduler, scaler, epoch, device=device)\n    val_loss, val_wer = validate(model, val_loader, epoch, device=device)\n\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        save_checkpoint(model, optimizer, scheduler, epoch, f\"/kaggle/working/model_best_epoch{epoch}.pt\")\n\nprint(\"Training complete ✅\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:15:31.171088Z","iopub.execute_input":"2025-11-07T05:15:31.171310Z","iopub.status.idle":"2025-11-07T05:34:13.978348Z","shell.execute_reply.started":"2025-11-07T05:15:31.171294Z","shell.execute_reply":"2025-11-07T05:34:13.977590Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**PHASE 5 OF THE PROJECT**","metadata":{}},{"cell_type":"code","source":"!ls -lh /kaggle/working | grep model_\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:34:13.979276Z","iopub.execute_input":"2025-11-07T05:34:13.979542Z","iopub.status.idle":"2025-11-07T05:34:14.173521Z","shell.execute_reply.started":"2025-11-07T05:34:13.979513Z","shell.execute_reply":"2025-11-07T05:34:14.172727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# Phase 5 - Cell 1 (updated)\n# Load the best available model checkpoint for inference\n# =========================================================\n\nimport torch\n\n# ✅ Use the best checkpoint that exists in /kaggle/working\nbest_ckpt_path = \"/kaggle/working/model_best_epoch4.pt\"\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Load checkpoint safely\nprint(f\"Loading checkpoint from: {best_ckpt_path}\")\nckpt = torch.load(best_ckpt_path, map_location=device)\n\n# Load model weights and set up for inference\nmodel.load_state_dict(ckpt[\"model_state_dict\"])\nmodel.to(device)\nmodel.eval()\n\nprint(f\"✅ Model loaded successfully from {best_ckpt_path}\")\nprint(f\"Device in use: {device}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:34:14.174828Z","iopub.execute_input":"2025-11-07T05:34:14.175481Z","iopub.status.idle":"2025-11-07T05:34:14.361180Z","shell.execute_reply.started":"2025-11-07T05:34:14.175449Z","shell.execute_reply":"2025-11-07T05:34:14.360522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# Phase 5 - Cell 2\n# Find all test .hdf5 files for inference\n# =========================================================\n\nfrom pathlib import Path\n\n# We already have INPUT_DIR defined from Phase 2\n# (the folder containing all hdf5_data_final/ files)\ntest_files = sorted([p for p in INPUT_DIR.rglob(\"data_test.hdf5\")])\n\nprint(f\"✅ Found {len(test_files)} test files.\\n\")\nprint(\"Example test files:\")\nfor f in test_files[:10]:\n    print(\"-\", f.relative_to(INPUT_DIR))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:34:14.361925Z","iopub.execute_input":"2025-11-07T05:34:14.362113Z","iopub.status.idle":"2025-11-07T05:34:14.562231Z","shell.execute_reply.started":"2025-11-07T05:34:14.362097Z","shell.execute_reply":"2025-11-07T05:34:14.561577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# Phase 5 - Cell 3\n# Helper functions for inference and greedy decoding\n# =========================================================\n\nimport torch\nimport numpy as np\nimport h5py\nfrom pathlib import Path\nfrom torch.amp import autocast\n\n@torch.no_grad()\ndef greedy_decode(log_probs):\n    \"\"\"\n    Performs greedy decoding of log probabilities.\n    Removes blanks (0) and repeated tokens.\n    Args:\n        log_probs: [B, T', V] tensor\n    Returns:\n        List of integer token sequences\n    \"\"\"\n    preds = log_probs.argmax(-1).cpu().numpy()  # [B, T']\n    decoded = []\n    for i in range(preds.shape[0]):\n        seq = preds[i]\n        seq_clean = [p for j, p in enumerate(seq)\n                     if (j == 0 or p != seq[j - 1]) and p != 0]\n        decoded.append(seq_clean)\n    return decoded\n\n\n@torch.no_grad()\ndef infer_file(model, path, device=\"cuda\"):\n    \"\"\"\n    Runs inference on all trials inside one HDF5 test file.\n    Args:\n        model: trained BrainToTextModel\n        path: path to .hdf5 file\n        device: 'cuda' or 'cpu'\n    Returns:\n        List of dicts: [{'id': '<filename>_<trial>', 'transcript': '<decoded tokens>'}, ...]\n    \"\"\"\n    preds_list = []\n    with h5py.File(path, \"r\") as hf:\n        for trial in hf.keys():\n            if \"input_features\" not in hf[trial]:\n                continue\n\n            feats = hf[trial][\"input_features\"][()].astype(\"float32\")\n            x = torch.from_numpy(feats).unsqueeze(0).permute(0, 2, 1).to(device)  # [1, C, T]\n\n            with autocast(\"cuda\", dtype=torch.float16):\n                log_probs = model(x)  # [1, T', V]\n\n            decoded = greedy_decode(log_probs)[0]\n            decoded_str = \" \".join(map(str, decoded))\n\n            preds_list.append({\n                \"id\": f\"{Path(path).stem}_{trial}\",\n                \"transcript\": decoded_str\n            })\n\n    return preds_list\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:34:14.563075Z","iopub.execute_input":"2025-11-07T05:34:14.563363Z","iopub.status.idle":"2025-11-07T05:34:14.572357Z","shell.execute_reply.started":"2025-11-07T05:34:14.563339Z","shell.execute_reply":"2025-11-07T05:34:14.571652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# Phase 5 - Cell 4\n# Run inference on all test files and collect predictions\n# =========================================================\n\nimport pandas as pd\nfrom tqdm import tqdm\n\nall_preds = []\n\nprint(\"🚀 Running inference on test set...\")\nfor test_path in tqdm(test_files, desc=\"Inference progress\"):\n    preds = infer_file(model, test_path, device=device)\n    all_preds.extend(preds)\n\n# Convert all predictions into a DataFrame\nsub_df = pd.DataFrame(all_preds)\nprint(f\"\\n✅ Inference complete. Total predictions: {len(sub_df)}\")\n\n# Show a small preview\ndisplay(sub_df.head(10))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:34:14.573150Z","iopub.execute_input":"2025-11-07T05:34:14.573467Z","iopub.status.idle":"2025-11-07T05:35:03.270141Z","shell.execute_reply.started":"2025-11-07T05:34:14.573438Z","shell.execute_reply":"2025-11-07T05:35:03.269529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# Phase 5 - Cell 5\n# Save predictions as submission.csv\n# =========================================================\nsub_path = \"/kaggle/working/submission.csv\"\nsub_df.to_csv(sub_path, index=False)\n\nprint(f\"✅ Submission file saved successfully at: {sub_path}\")\nprint(\"Preview:\")\ndisplay(sub_df.head(10))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:35:03.270841Z","iopub.execute_input":"2025-11-07T05:35:03.271092Z","iopub.status.idle":"2025-11-07T05:35:03.286365Z","shell.execute_reply.started":"2025-11-07T05:35:03.271065Z","shell.execute_reply":"2025-11-07T05:35:03.285661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# Phase 6: Show sample predictions (actual vs predicted)\n# =========================================================\nimport random\nfrom torch.amp import autocast\n\n# Pick a few random validation files\nval_files = sorted([p for p in INPUT_DIR.rglob(\"data_val.hdf5\")])\nprint(f\"Found {len(val_files)} validation files.\")\nsample_files = random.sample(val_files, 3)  # show samples from 3 val files\n\nfor val_path in sample_files:\n    print(f\"\\n📘 File: {val_path.name}\")\n    with h5py.File(val_path, \"r\") as hf:\n        # Pick 2 random trials from each file\n        trial_names = random.sample(list(hf.keys()), 2)\n        for trial in trial_names:\n            if \"input_features\" not in hf[trial]:\n                continue\n            feats = hf[trial][\"input_features\"][()].astype(\"float32\")\n            transcript = hf[trial][\"transcription\"][()].astype(\"int32\")\n\n            # Decode transcript (true)\n            true_text = \" \".join(map(str, transcript.tolist()))\n\n            # Model prediction\n            x = torch.from_numpy(feats).unsqueeze(0).permute(0, 2, 1).to(device)\n            with autocast(\"cuda\", dtype=torch.float16):\n                log_probs = model(x)\n            decoded_seq = greedy_decode(log_probs)[0]\n            pred_text = \" \".join(map(str, decoded_seq))\n\n            print(f\"🧠 Trial: {trial}\")\n            print(f\"   ➤ Predicted: {pred_text[:200]}...\")\n            print(f\"   ➤ Actual:    {true_text[:200]}...\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:35:03.287117Z","iopub.execute_input":"2025-11-07T05:35:03.287436Z","iopub.status.idle":"2025-11-07T05:35:03.799806Z","shell.execute_reply.started":"2025-11-07T05:35:03.287419Z","shell.execute_reply":"2025-11-07T05:35:03.799016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# Convert integer predictions & actuals to readable text\n# =========================================================\n\ndef decode_ascii(int_seq):\n    \"\"\"Convert list/array of integers to readable text (ASCII decoding).\"\"\"\n    chars = [chr(i) for i in int_seq if 32 <= i <= 126]  # printable characters\n    return \"\".join(chars)\n\n# Display a few samples again in readable form\nval_path = random.choice(val_files)\nprint(f\"\\n📘 File: {val_path.name}\")\n\nwith h5py.File(val_path, \"r\") as hf:\n    trial_names = random.sample(list(hf.keys()), 3)\n    for trial in trial_names:\n        if \"input_features\" not in hf[trial]:\n            continue\n\n        feats = hf[trial][\"input_features\"][()].astype(\"float32\")\n        transcript = hf[trial][\"transcription\"][()].astype(\"int32\")\n\n        # Run prediction\n        x = torch.from_numpy(feats).unsqueeze(0).permute(0, 2, 1).to(device)\n        with autocast(\"cuda\", dtype=torch.float16):\n            log_probs = model(x)\n        decoded_seq = greedy_decode(log_probs)[0]\n\n        # Decode to readable text\n        pred_text = decode_ascii(decoded_seq)\n        true_text = decode_ascii(transcript.tolist())\n\n        print(f\"🧠 Trial: {trial}\")\n        print(f\"   ➤ Predicted: {pred_text}\")\n        print(f\"   ➤ Actual:    {true_text}\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-07T05:35:03.800617Z","iopub.execute_input":"2025-11-07T05:35:03.800893Z","iopub.status.idle":"2025-11-07T05:35:03.947620Z","shell.execute_reply.started":"2025-11-07T05:35:03.800857Z","shell.execute_reply":"2025-11-07T05:35:03.947028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}