{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"},{"sourceId":14328691,"sourceType":"datasetVersion","datasetId":9147692},{"sourceId":289010823,"sourceType":"kernelVersion"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:01.577465Z","iopub.execute_input":"2025-12-30T07:27:01.577695Z","iopub.status.idle":"2025-12-30T07:27:01.892337Z","shell.execute_reply.started":"2025-12-30T07:27:01.577677Z","shell.execute_reply":"2025-12-30T07:27:01.891622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import List, Optional\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport math","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:01.893527Z","iopub.execute_input":"2025-12-30T07:27:01.893926Z","iopub.status.idle":"2025-12-30T07:27:05.315457Z","shell.execute_reply.started":"2025-12-30T07:27:01.893909Z","shell.execute_reply":"2025-12-30T07:27:05.314626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndevice0 = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndevice1 = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# if xm.xla_device(n=0, dev_type='TPU') is not None:\n#     device = xm.xla_device()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.316297Z","iopub.execute_input":"2025-12-30T07:27:05.316610Z","iopub.status.idle":"2025-12-30T07:27:05.367524Z","shell.execute_reply.started":"2025-12-30T07:27:05.316592Z","shell.execute_reply":"2025-12-30T07:27:05.366542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MLP(nn.Module):\n    \"\"\"Multi-layer perceptron (MLP) model.\"\"\"\n\n    def __init__(\n        self,\n        input_size: int,\n        hidden_size: int,\n        num_output: int,\n        num_layers: int,\n    ):\n        super(MLP, self).__init__()\n        next_input_size = input_size\n\n        layers = []\n        if num_layers>1:\n            for _ in range(num_layers):\n                layers.append(nn.Linear(next_input_size, hidden_size))\n                layers.append(nn.ReLU())\n                next_input_size = hidden_size\n\n        layers.append(nn.Linear(next_input_size, num_output))\n        self.network = nn.Sequential(*layers)\n\n    def forward(self, x):\n        return self.network(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.369022Z","iopub.execute_input":"2025-12-30T07:27:05.369259Z","iopub.status.idle":"2025-12-30T07:27:05.509234Z","shell.execute_reply.started":"2025-12-30T07:27:05.369241Z","shell.execute_reply":"2025-12-30T07:27:05.508582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def inverse_scaled_sigmoid(x, lower_bound: float, upper_bound: float):\n    x = torch.clamp(x, lower_bound + 1e-6, upper_bound - 1e-6)\n    return torch.log((x - lower_bound) / (upper_bound - x))\ndef custom_tanh(x):\n    return torch.tanh(x * 2 / 3) * 1.7159\n\ndef create_interlocking_indices(num_input: int):\n    half_num_input_data = num_input // 2\n    half_range_steps = (torch.arange(num_input) % 2) * half_num_input_data\n    single_steps = torch.div(torch.arange(num_input), 2, rounding_mode=\"floor\")\n    return half_range_steps + single_steps\n\ndef create_overlapping_window_indices(\n    num_input: int, num_windows: int, num_elements_per_window: int\n):\n    stride_size = math.ceil(num_input / num_windows)\n    overlapping_indices = (\n        torch.arange(num_windows).unsqueeze(1) * stride_size\n    ) + torch.arange(num_elements_per_window).unsqueeze(0)\n    valid_indices = overlapping_indices < num_input\n    overlapping_indices = torch.clamp(overlapping_indices, max=num_input - 1)  # fix\n    return overlapping_indices.flatten(), valid_indices.flatten()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.509889Z","iopub.execute_input":"2025-12-30T07:27:05.510128Z","iopub.status.idle":"2025-12-30T07:27:05.521541Z","shell.execute_reply.started":"2025-12-30T07:27:05.510082Z","shell.execute_reply":"2025-12-30T07:27:05.520790Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import h5py\n\ndef load_h5py_file(file_path):\n    data = {\n        'neural_features': [],\n        'n_time_steps': [],\n        'seq_class_ids': [],\n        'seq_len': [],\n        'transcriptions': [],\n        'sentence_label': [],\n        'session': [],\n        'block_num': [],\n        'trial_num': [],\n    }\n    # Open the hdf5 file for that day\n    with h5py.File(file_path, 'r') as f:\n\n        keys = list(f.keys())\n\n        # For each trial in the selected trials in that day\n        for key in keys:\n            g = f[key]\n\n            neural_features = g['input_features'][:]\n            n_time_steps = g.attrs['n_time_steps']\n            seq_class_ids = g['seq_class_ids'][:] if 'seq_class_ids' in g else None\n            seq_len = g.attrs['seq_len'] if 'seq_len' in g.attrs else None\n            transcription = g['transcription'][:] if 'transcription' in g else None\n            sentence_label = g.attrs['sentence_label'][:] if 'sentence_label' in g.attrs else None\n            session = g.attrs['session']\n            block_num = g.attrs['block_num']\n            trial_num = g.attrs['trial_num']\n\n            data['neural_features'].append(neural_features)\n            data['n_time_steps'].append(n_time_steps)\n            data['seq_class_ids'].append(seq_class_ids)\n            data['seq_len'].append(seq_len)\n            data['transcriptions'].append(transcription)\n            data['sentence_label'].append(sentence_label)\n            data['session'].append(session)\n            data['block_num'].append(block_num)\n            data['trial_num'].append(trial_num)\n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.522310Z","iopub.execute_input":"2025-12-30T07:27:05.522550Z","iopub.status.idle":"2025-12-30T07:27:05.675684Z","shell.execute_reply.started":"2025-12-30T07:27:05.522534Z","shell.execute_reply":"2025-12-30T07:27:05.675040Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PREPROCESS_CONFIGURATIONS = [None, \"random_routing\", \"neuronio_routing\"]\n\n\nclass ELM(nn.Module):\n    \"\"\"Expressive Leaky Memory (ELM) neuron model.\"\"\"\n\n    __constants__ = [\n        \"num_input\",\n        \"num_output\",\n        \"num_memory\",\n        \"lambda_value\",\n        \"mlp_num_layers\",\n        \"mlp_activation\",\n        \"memory_tau_min\",\n        \"memory_tau_max\",\n        \"learn_memory_tau\",\n        \"w_s_value\",\n        \"num_synapse_per_branch\",\n        \"input_to_synapse_routing\",\n        \"delta_t\",\n    ]\n\n    def __init__(\n        self,\n        num_input: int,\n        num_output: int,\n        num_memory: int = 100,\n        lambda_value: float = 5.0,\n        mlp_num_layers: int = 1,\n        mlp_hidden_size: Optional[int] = None,\n        mlp_activation: str = \"relu\",\n        mlp_grad: bool = True,\n        tau_s_value: float = 5.0,\n        memory_tau_min: float = 1.0,\n        memory_tau_max: float = 1000.0,\n        learn_memory_tau: bool = False,\n        w_s_value: float = 0.5,\n        num_branch: Optional[int] = None,\n        num_synapse_per_branch: int = 1,\n        input_to_synapse_routing: Optional[str] = None,\n        delta_t: float = 1.0,\n    ):\n        super(ELM, self).__init__()\n        self.num_input, self.num_output = num_input, num_output\n        self.num_memory = num_memory\n        self.lambda_value = lambda_value\n        self.mlp_num_layers = mlp_num_layers\n        self.mlp_activation = mlp_activation\n        self.memory_tau_min, self.memory_tau_max = memory_tau_min, memory_tau_max\n        self.learn_memory_tau = learn_memory_tau\n        self.tau_s_value, self.w_s_value = tau_s_value, w_s_value\n        self.num_synapse_per_branch = num_synapse_per_branch\n        self.input_to_synapse_routing = input_to_synapse_routing\n        self.delta_t = delta_t\n\n        # derived neuron properties\n        self.mlp_hidden_size = mlp_hidden_size if mlp_hidden_size else 2 * num_memory\n        self.num_branch = self.num_input if num_branch is None else num_branch\n        self.num_mlp_input = self.num_branch + num_memory\n        self.num_synapse = num_synapse_per_branch * self.num_branch\n\n        # sanity check of input configuration\n        assert self.num_synapse == num_input or input_to_synapse_routing is not None\n        assert self.input_to_synapse_routing in PREPROCESS_CONFIGURATIONS\n\n        # initialization of model weights\n        self.mlp = MLP(\n            self.num_mlp_input,\n            self.mlp_hidden_size,\n            num_memory,\n            mlp_num_layers,\n        )\n\n        if not mlp_grad:\n            for param in self.mlp.parameters():\n                param.requires_grad = False\n        \n        self._proto_w_s = nn.parameter.Parameter(\n            torch.full((self.num_synapse,), w_s_value)\n        )\n        self.w_y = nn.Linear(num_memory, num_output)\n\n        # initialization of synapse time constants and decay factors\n        tau_s = torch.full((self.num_synapse,), tau_s_value)\n        self.tau_s = nn.parameter.Parameter(tau_s, requires_grad=False)\n\n        # initialization of memory time constants and decay factors\n        _proto_tau_m = torch.logspace(\n            math.log10(memory_tau_min + 1e-6),\n            math.log10(memory_tau_max - 1e-6),\n            num_memory,\n        )\n        _proto_tau_m = inverse_scaled_sigmoid(\n            _proto_tau_m, memory_tau_min, memory_tau_max\n        )\n        self._proto_tau_m = nn.parameter.Parameter(\n            _proto_tau_m, requires_grad=learn_memory_tau\n        )\n\n        # NOTE: part of model for ease of use\n        routing_artifacts = self.create_input_to_synapse_indices()\n        self.input_to_synapse_indices = nn.parameter.Parameter(\n            routing_artifacts[0], requires_grad=False\n        )\n        self.valid_indices_mask = nn.parameter.Parameter(\n            routing_artifacts[1], requires_grad=False\n        )\n\n        # self.in_norm = nn.LayerNorm(self.num_input)\n\n    @property\n    def tau_m(self):\n        return (self.memory_tau_max - self.memory_tau_min) * \\\n        torch.sigmoid(self._proto_tau_m) + self.memory_tau_min\n\n    @property\n    def kappa_m(self):\n        return torch.exp(-self.delta_t / torch.clamp(self.tau_m, min=1e-6))\n\n    @property\n    def kappa_s(self):\n        return torch.exp(-self.delta_t / torch.clamp(self.tau_s, min=1e-6))\n\n    @property\n    def w_s(self):\n        return torch.relu(self._proto_w_s)\n\n    # NOTE: part of model for ease of use\n    def create_input_to_synapse_indices(self):\n        if self.input_to_synapse_routing == \"random_routing\":\n            # randomly select num_synapse from num_input\n            input_to_synapse_indices = torch.randint(\n                self.num_input, (self.num_synapse,)\n            )\n            return input_to_synapse_indices, torch.ones_like(input_to_synapse_indices)\n        elif self.input_to_synapse_routing == \"neuronio_routing\":\n            # sanity check of input configuration\n            assert (\n                math.ceil(self.num_input / self.num_branch)\n                <= self.num_synapse_per_branch\n            )\n\n            # interlace excitatory and inhibitory inputs\n            interlocking_indices = create_interlocking_indices(self.num_input)\n            # assign neighbouring inputs to same branch\n            overlapping_indices, valid_indices_mask = create_overlapping_window_indices(\n                self.num_input, self.num_branch, self.num_synapse_per_branch\n            )\n            input_to_synapse_indices = interlocking_indices[overlapping_indices]\n\n            return input_to_synapse_indices, valid_indices_mask\n        else:\n            return None, None\n\n    # NOTE: part of model for ease of use\n    def route_input_to_synapses(self, x):\n        if self.input_to_synapse_routing is not None:\n            x = torch.index_select(x, 2, self.input_to_synapse_indices)\n            x = x * self.valid_indices_mask  # valid mask\n        return x\n\n    def dynamics(self, x, s_prev, m_prev, w_s, kappa_s, kappa_m):\n        # compute the dynamics for a single timestep\n        batch_size, _ = x.shape\n        s_t = kappa_s * s_prev + w_s * x\n        syn_input = s_t.view(batch_size, self.num_branch, -1).sum(dim=-1)\n        delta_m_t = custom_tanh(\n            self.mlp(torch.cat([syn_input, kappa_m * m_prev], dim=-1))\n        )\n        m_t = kappa_m * m_prev + self.lambda_value * (1 - kappa_m) * delta_m_t\n        y_t = self.w_y(m_t)\n        return y_t, s_t, m_t\n\n    def forward(self, X):\n        # compute the the recurrent dynamics for a sample\n        # X = self.in_norm(X)\n        batch_size, T, _ = X.shape\n        w_s = self.w_s\n        kappa_s, kappa_m = self.kappa_s, self.kappa_m\n        s_prev = torch.zeros(batch_size, len(kappa_s), device=X.device)\n        m_prev = torch.zeros(batch_size, len(kappa_m), device=X.device)\n        outputs = torch.jit.annotate(List[torch.Tensor], [])\n        inputs = self.route_input_to_synapses(X)\n        for t in range(T):\n            y_t, s_prev, m_prev = self.dynamics(\n                inputs[:, t], s_prev, m_prev, w_s, kappa_s, kappa_m\n            )\n            outputs.append(y_t)\n        output = torch.stack(outputs, dim=-1)\n        # output = self.mean(output)\n        return output.permute(0,2,1) # Output in shape of [B, T, F]\n        return s_prev, m_prev\n\n    # NOTE: part of model for ease of use\n    def neuronio_eval_forward(\n        self, X, y_train_soma_scale: float = 1/10\n    ):\n        outputs = self.forward(X)\n        spike_pred, soma_pred = outputs[..., 0], outputs[..., 1]\n\n        # apply sigmoid to spike (probability) prediction\n        spike_pred = torch.sigmoid(spike_pred)\n        # apply soma scale to soma prediction\n        soma_pred = 1 / y_train_soma_scale * soma_pred\n\n        return torch.stack([spike_pred, soma_pred], dim=-1)\n\n    # NOTE: part of model for ease of use\n    def neuronio_viz_forward(\n        self, X, y_train_soma_scale: float = 1/10\n    ):\n        # compute the the recurrent dynamics for a sample\n        batch_size, T, _ = X.shape\n        w_s = self.w_s\n        kappa_s, kappa_m = self.kappa_s, self.kappa_m\n        s_prev = torch.zeros(batch_size, len(kappa_s), device=X.device)\n        m_prev = torch.zeros(batch_size, len(kappa_m), device=X.device)\n\n        # calcualte the outputs, synapse and memory values\n        outputs = torch.jit.annotate(List[torch.Tensor], [])\n        s_record = torch.jit.annotate(List[torch.Tensor], [])\n        m_record = torch.jit.annotate(List[torch.Tensor], [])\n        inputs = self.route_input_to_synapses(X)\n        for t in range(T):\n            y_t, s_prev, m_prev = self.dynamics(\n                inputs[:, t], s_prev, m_prev, w_s, kappa_s, kappa_m\n            )\n            outputs.append(y_t)\n            s_record.append(s_prev)\n            m_record.append(m_prev)\n        outputs = torch.stack(outputs, dim=-2)\n        s_record = torch.stack(s_record, dim=-2)\n        m_record = torch.stack(m_record, dim=-2)\n\n        # postprocess the outputs\n        spike_pred, soma_pred = outputs[..., 0], outputs[..., 1]\n        spike_pred = torch.sigmoid(spike_pred)\n        soma_pred = 1 / y_train_soma_scale * soma_pred\n        outputs = torch.stack([spike_pred, soma_pred], dim=-1)\n\n        return outputs, s_record, m_record","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.676426Z","iopub.execute_input":"2025-12-30T07:27:05.676748Z","iopub.status.idle":"2025-12-30T07:27:05.703520Z","shell.execute_reply.started":"2025-12-30T07:27:05.676730Z","shell.execute_reply":"2025-12-30T07:27:05.702873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nfrom collections import defaultdict\nimport itertools","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.704274Z","iopub.execute_input":"2025-12-30T07:27:05.704580Z","iopub.status.idle":"2025-12-30T07:27:05.723808Z","shell.execute_reply.started":"2025-12-30T07:27:05.704525Z","shell.execute_reply":"2025-12-30T07:27:05.722876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def find_factor_pairs(n):\n    \"\"\"\n    Finds all factor pairs of a given positive integer n.\n    Factors appear in pairs (i, n//i).\n    \"\"\"\n    pairs = []\n    # Loop up to the square root of n to find all pairs efficiently.\n    for i in range(1, int(math.sqrt(n)) + 1):\n        if n % i == 0:\n            # If i divides n evenly, add the pair (i, n//i) to the list.\n            pair = (i, n // i)\n            pairs.append(pair)\n    return pairs\n\ndef choose_random_pair(n):\n    \"\"\"\n    Randomly selects one pair from the list of factor pairs.\n    \"\"\"\n    pair_list = find_factor_pairs(n)\n    if not pair_list:\n        return None\n    # print(pair_list)\n    # Use random.choice to select a single random item from the list.\n    # print(pair_list)\n    return pair_list[-1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.725062Z","iopub.execute_input":"2025-12-30T07:27:05.725343Z","iopub.status.idle":"2025-12-30T07:27:05.738836Z","shell.execute_reply.started":"2025-12-30T07:27:05.725316Z","shell.execute_reply":"2025-12-30T07:27:05.738093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def xavier_weights(m):\n    if isinstance(m, nn.Linear):\n        # Apply Xavier uniform initialization to weights\n        nn.init.xavier_uniform_(m.weight)\n        # Initialize biases to zero if they exist\n        if m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n\ndef kaiming_weights(m):\n    if isinstance(m, nn.Linear):\n        # Apply Xavier uniform initialization to weights\n        nn.init.kaiming_uniform_(m.weight)\n        # Initialize biases to zero if they exist\n        if m.bias is not None:\n            nn.init.constant_(m.bias, 0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.742210Z","iopub.execute_input":"2025-12-30T07:27:05.742543Z","iopub.status.idle":"2025-12-30T07:27:05.754356Z","shell.execute_reply.started":"2025-12-30T07:27:05.742521Z","shell.execute_reply":"2025-12-30T07:27:05.753616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Node(nn.Module):\n    def __init__(self, params: dict, node_index, output_index, input_index: int = 0):\n        super(Node, self).__init__()\n\n        self.connected = False\n        self.node_index = node_index\n        self.input_nodes = set([input_index])\n        self.parents_set = set([input_index])\n        self.output_index= set([output_index])\n        self.output_nodes = set([])\n        self._threshold = nn.Parameter(torch.tensor(1.0))\n        self.params = params.copy()\n        self.params['mlp_grad'] = True\n        self.initialized = False\n        self.root_node = True\n        self.device = None\n\n    @property\n    def threshold(self):\n        return torch.tanh(self._threshold)\n\n    def connect(self, node, readout_index):\n\n        if -1 in self.input_nodes: # Remove the Root input if taking input from another nodes\n            self.input_nodes.remove(-1)\n        \n        self.root_node = False # Node do not takes input direct from real input X\n\n        self.parents_set = self.parents_set | node.parents_set | {node.node_index}\n        self.input_nodes.add(node)\n        # self.input_index = node.node_index\n\n        node.output_nodes.add(self)\n        \n        \n        self.params['num_input'] = sum([node.params['num_output'] for node in self.input_nodes])\n        if self.params['num_output']!=1: \n            self.params['num_output'] = max(node.params['num_output']//4, 2)\n        # self.params['mlp_grad'] = False\n\n    def initialize(self):\n\n        # synaptic_pair = choose_random_pair(self.params['num_input'])\n        if not self.root_node:\n            self.params[\"input_to_synapse_routing\"] = None\n            self.params[\"num_branch\"] = self.params['num_input']\n            self.params[\"num_synapse_per_branch\"] = 1\n        \n\n        if self.root_node:\n            self.device=device0\n        self.node = ELM(**self.params).to(self.device)\n        \n        self.node.apply(xavier_weights)\n\n        self.initialized = True\n        # print(\"initialized\")\n    \n    def forward(self, x):\n\n        if not self.initialized:\n            self.initialize()\n        if x.device != self.device:           \n            x = x.to(self.device)\n        output = self.node(x)\n        # print(self.threshold)\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.755519Z","iopub.execute_input":"2025-12-30T07:27:05.755994Z","iopub.status.idle":"2025-12-30T07:27:05.770837Z","shell.execute_reply.started":"2025-12-30T07:27:05.755955Z","shell.execute_reply":"2025-12-30T07:27:05.770034Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Reservoir(nn.Module):\n    def __init__(\n        self,\n        internal_nodes_params, # Dict for internal nodes param\n        internal_output_dim, # Number of output per internal nodes\n        num_internal_nodes = 5, # Number of internal nodes\n        internal_connectivity = 0.2,\n    ):\n        super(Reservoir, self).__init__()\n        self.internal_nodes = nn.ModuleList([])\n        self.internal_connectivity = internal_connectivity\n\n        all_indices = list(range(num_internal_nodes))\n        \n        # all_indices.append(len(all_indices)+1)\n        # all_indices.insert(0, 0)\n        self.internal_nodes_index = {}\n        self.connectivity = defaultdict(list)\n        device_a = True\n        for node_index in all_indices:\n            self.internal_nodes.append(Node(params=internal_nodes_params,\n                                            node_index=node_index,\n                                            input_index=-1,\n                                            output_index=num_internal_nodes+1\n                                           ))\n            self.connectivity[node_index].append(num_internal_nodes+1)\n            if device_a:\n                self.internal_nodes[node_index].device = device0\n            else:\n                self.internal_nodes[node_index].device = device1\n            device_a = not device_a \n        if internal_connectivity < 1:\n            pairs = list(itertools.combinations(all_indices, 2))\n            for (k, l) in pairs:\n                connect = random.random() < internal_connectivity\n                if not connect:\n                    continue\n                parent = self.internal_nodes[k]\n                child  = self.internal_nodes[l]\n    \n                # cycle check using ancestry set\n                if (\n                    parent.node_index in child.parents_set\n                    # or -1 not in child.input_nodes\n                    # or any(child.node_index in children for children in self.connectivity.values())\n                ):\n                    continue\n                self.connectivity[parent.node_index].append(child.node_index)\n        else:\n            for i in range(num_internal_nodes):\n                j = i+1\n                if j < num_internal_nodes:\n                    parent = self.internal_nodes[i]\n                    child = self.internal_nodes[j]\n                    self.connectivity[parent.node_index].append(child.node_index)\n        \n        self.exec_order = []\n        self.initialize()\n        \n    def dfs(self, node, visited):\n        if node.node_index in visited or len(self.exec_order) == len(self.internal_nodes):\n            return visited\n        parents = node.input_nodes\n        if -1 not in parents:\n            for parent in parents:\n                self.dfs(parent, visited)\n        self.exec_order.append(node)\n        visited.add(node.node_index)\n        return visited\n            \n\n    def initialize(self, output_dim=None):\n        self.num_input = 0\n\n        # Build Graph Structure\n        for p, c in self.connectivity.items():\n            root_index = p\n            for index in c:\n                if index > len(self.internal_nodes):\n                    continue\n                internal_inputs = index\n                node = self.internal_nodes[internal_inputs]\n                root_node = self.internal_nodes[root_index]\n                \n                node.connect(root_node, len(self.internal_nodes)+1)\n\n        # Building Execution Order\n        visited = set()\n        for node in self.internal_nodes:\n            visited = self.dfs(node, visited)\n\n    def save_model(self):\n        checkpoint = {\n            'state': self.state_dict(),\n            'internal_node_params': [node.params for node in self.internal_nodes],\n            'connectivity': self.connectivity,\n            'exec_order': self.exec_order\n        }\n        return checkpoint\n\n    def load_model(self, checkpoint, strict=True):\n        # 1. Restore graph structure\n        self.connectivity = checkpoint['connectivity']\n        \n        for i, node in enumerate(self.internal_nodes):\n            node.params = checkpoint['internal_node_params'][i]\n            \n        for node in self.internal_nodes:\n            node.initialize()\n        self.exec_order = checkpoint['exec_order']\n    \n        # 2. Load parameters\n        self.load_state_dict(checkpoint['state'], strict=strict)\n\n    \n    def forward(self, x):\n        outputs = [None] * len(self.internal_nodes)\n        for i, node in enumerate(self.exec_order):\n            if -1 in node.input_nodes:\n                inp = x\n            else:\n                parents = node.input_nodes\n                inp = []\n                for parent in parents:\n                    inp.append(outputs[parent.node_index])\n                inp = torch.cat(inp, dim=-1) # All tensors in in device0 it'll work with any problem\n            output, state = node(inp)\n            print('state shape:', state.shape)\n            if output.device != device0:\n                output = output.to(device0) # Move to device0 if not already in 0\n            outputs[node.node_index] = output\n        \n        final_output = []\n        for output in outputs:\n            if output.device != device0: # Just to be sure\n                output = output.to(device0)\n            final_output.append(output)\n        del outputs\n        final_output = torch.cat(final_output, dim=-1)\n\n        return final_output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.771697Z","iopub.execute_input":"2025-12-30T07:27:05.771924Z","iopub.status.idle":"2025-12-30T07:27:05.791577Z","shell.execute_reply.started":"2025-12-30T07:27:05.771903Z","shell.execute_reply":"2025-12-30T07:27:05.790926Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"kernel_sizes = [7, 17, 29, 65]\n\nclass CResT(nn.Module):\n    def __init__(\n        self,\n        num_input,\n        num_output,\n        reservoir_params,\n        conv_feature_out=256,\n        readout_mlp_layers=1,\n        mlp_hidden_size=None\n    ):\n        super(CResT, self).__init__()\n\n        self.num_output = num_output\n        self.readout_mlp_layers = readout_mlp_layers\n        \n        self.conv = nn.Sequential(\n                        # Local features\n                        nn.Conv1d(in_channels=num_input, out_channels=256,\n                                  kernel_size=5, padding=2, stride=2),\n                        nn.GELU(),\n                    \n                        # Mid-range\n                        nn.Conv1d(256, 256, kernel_size=9, padding=4, stride=2),\n                        nn.GELU(),\n                    \n                        # Long context via dilation\n                        nn.Conv1d(256, 256, kernel_size=7, padding=6, stride=2, dilation=2),\n                        nn.GELU(),\n                    \n                        # nn.LayerNorm(256),\n                    \n                        # Optional projection to reservoir dims\n                        nn.Conv1d(256, conv_feature_out, kernel_size=1)\n                    )\n\n    \n        self.reservoir = ELM(**reservoir_params)\n        self.reservoir_output = reservoir_params['num_output']\n\n        self.sequeezer = nn.ModuleList([\n            nn.Conv1d(self.reservoir_output, 10, kernel_size=k, \n                      padding=(k-1)//2, stride=2) for k in kernel_sizes\n        ])\n\n        # self.mixer = nn.Conv1d(10*len(kernel_sizes), )\n\n        self.mlp_hidden_size = mlp_hidden_size if mlp_hidden_size else self.reservoir_output // 2\n\n        \n        \n        self.readout_norm = nn.LayerNorm(40)\n        self.readout = nn.Conv1d(\n            in_channels=40, out_channels=self.num_output, kernel_size=1)\n\n        self.conv.apply(kaiming_weights)\n        self.reservoir.apply(xavier_weights)\n        self.readout.apply(kaiming_weights)\n        \n    def save_model(self):\n        checkpoint = {\n            'conv_state': self.conv.state_dict(),\n            'readout_state': self.readout.state_dict(),\n            'reservoir_state': self.reservoir.state_dict(),\n            'sequeezer': self.sequeezer.state_dict()\n        }\n        return checkpoint\n\n    def load_model(self, checkpoint):\n        self.conv.load_state_dict(checkpoint['conv_state'])\n        self.reservoir.load_state_dict(checkpoint['reservoir_state'])\n        self.readout.load_state_dict(checkpoint['readout_state'])\n        self.sequeezer.load_state_dict(checkpoint['sequeezer_state'])\n    \n    def forward(self, x):\n\n        x = x.permute(0, 2, 1)    # (B, input_dim, T)\n        conv_out = self.conv(x)\n        # print('x', x.shape)\n        # print('conv', conv_out.shape)\n        conv_out = conv_out.permute(0, 2, 1)    # (B, input_dim, T)\n        reservoir_features = self.reservoir(conv_out)\n        squeezed = []\n        reservoir_features = reservoir_features.permute(0, 2, 1)    # (B, input_dim, T)\n        for s in self.sequeezer:\n            out = s(reservoir_features).permute(0, 2, 1)\n            squeezed.append(out)\n        padded_squeezed = torch.cat(squeezed, dim=-1)\n        \n        padded_squeezed = self.readout_norm(padded_squeezed).permute(0, 2, 1) \n        output = self.readout(padded_squeezed)\n        \n        return output.permute(0, 2, 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:58.004863Z","iopub.execute_input":"2025-12-30T07:27:58.005223Z","iopub.status.idle":"2025-12-30T07:27:58.016636Z","shell.execute_reply.started":"2025-12-30T07:27:58.005197Z","shell.execute_reply":"2025-12-30T07:27:58.015816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LOGIT_TO_PHONEME = [\n'BLANK',    # \"BLANK\" = CTC blank symbol\n'AA', 'AE', 'AH', 'AO', 'AW',\n'AY', 'B', 'CH', 'D', 'DH',\n'EH', 'ER', 'EY', 'F', 'G',\n'HH', 'IH', 'IY', 'JH', 'K',\n'L', 'M', 'N', 'NG', 'OW',\n'OY', 'P', 'R', 'S', 'SH',\n'T', 'TH', 'UH', 'UW', 'V',\n'W', 'Y', 'Z', 'ZH',\n' | ',    # \"|\" = silence token\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.808177Z","iopub.execute_input":"2025-12-30T07:27:05.808368Z","iopub.status.idle":"2025-12-30T07:27:05.821990Z","shell.execute_reply.started":"2025-12-30T07:27:05.808345Z","shell.execute_reply":"2025-12-30T07:27:05.821446Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LOGIT_TO_PHONEME[40]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.822713Z","iopub.execute_input":"2025-12-30T07:27:05.823264Z","iopub.status.idle":"2025-12-30T07:27:05.835647Z","shell.execute_reply.started":"2025-12-30T07:27:05.823247Z","shell.execute_reply":"2025-12-30T07:27:05.835036Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\ndef visualize_graph_flow(\n    connections,\n    input_index=None,\n    readout_index=None,\n    seed=42\n):\n    \"\"\"\n    connections: dict {node: set(children)}\n    input_index: node representing input (placed leftmost)\n    readout_index: node representing readout/output (placed rightmost)\n    \"\"\"\n\n    rng = np.random.default_rng(seed)\n    nodes = list(connections.keys())\n\n    # Separate node groups\n    intermediate_nodes = [\n        n for n in nodes\n        if n not in {input_index, readout_index}\n    ]\n\n    positions = {}\n\n    # --- Fixed positions ---\n    if input_index is not None:\n        positions[input_index] = (0.0, 0.0)\n\n    if readout_index is not None:\n        positions[readout_index] = (1.0, 0.0)\n\n    # --- Scatter intermediate nodes between them ---\n    for n in intermediate_nodes:\n        x = rng.uniform(0.15, 0.85)\n        y = rng.uniform(-0.5, 0.5)\n        positions[n] = (x, y)\n\n    # --- Plot ---\n    fig, ax = plt.subplots(figsize=(14, 8))\n    ax.set_aspect('equal')\n\n    # Draw nodes\n    for node, (x, y) in positions.items():\n        if node == input_index:\n            color = \"lightgreen\"\n        elif node == readout_index:\n            color = \"salmon\"\n        else:\n            color = \"skyblue\"\n\n        ax.scatter(x, y, s=300, color=color, edgecolor=\"black\", zorder=3)\n        ax.text(x, y, str(node), ha=\"center\", va=\"center\", fontsize=10, zorder=4)\n\n    # Draw directed edges\n    for src, children in connections.items():\n        if src not in positions:\n            continue\n\n        x1, y1 = positions[src]\n\n        for dst in children:\n            if dst not in positions:\n                continue\n\n            x2, y2 = positions[dst]\n\n            ax.annotate(\n                \"\",\n                xy=(x2, y2),\n                xytext=(x1, y1),\n                arrowprops=dict(\n                    arrowstyle=\"->\",\n                    color=\"black\",\n                    lw=1.2,\n                    alpha=0.8\n                ),\n                zorder=3\n            )\n\n    # Clean look\n    ax.set_xlim(-0.1, 1.1)\n    ax.set_ylim(-0.7, 0.7)\n    ax.set_xticks([])\n    ax.set_yticks([])\n    ax.set_title(\"Reservoir Graph – Input → Reservoir → Readout\")\n\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.836414Z","iopub.execute_input":"2025-12-30T07:27:05.837067Z","iopub.status.idle":"2025-12-30T07:27:05.848782Z","shell.execute_reply.started":"2025-12-30T07:27:05.837042Z","shell.execute_reply":"2025-12-30T07:27:05.848272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import ConcatDataset, Dataset\n\ndef temporal_mask(data, mask_percentage=0.05, mask_value=0.0):\n    \"\"\"\n    Applies temporal masking to a 2D tensor [Sequence, Features].\n    \"\"\"\n    if not torch.is_tensor(data):\n        data = torch.tensor(data, dtype=torch.float32)\n        \n    seq_len, _ = data.shape\n    num_to_mask = int(seq_len * mask_percentage)\n    \n    if num_to_mask > 0:\n        mask_indices = torch.randperm(seq_len)[:num_to_mask]\n        data[mask_indices, :] = mask_value\n        \n    return data\n\nclass BrainDataset(Dataset):\n    \"\"\"\n    Reads data from a single HDF5 file (e.g., data_train.hdf5).\n    - input_key: The name of the HDF5 dataset for input features.\n    - target_key: The name of the HDF5 dataset for target sequences (indices).\n    - is_test: If True, __getitem__ also returns the trial_key.\n    \"\"\"\n    def __init__(self, hdf5_file, input_key=\"input_features\", target_key=\"phoneme_indices\", is_test=False, use_augmentation=False):\n        self.file_path = hdf5_file\n        self.input_key = input_key\n        self.target_key = target_key # Key for sequence targets\n        self.is_test = is_test\n        \n        # --- FIX 2: STORE THE PARAMETER ---\n        self.use_augmentation = use_augmentation \n        self.file = None # File handle\n        \n        try:\n            with h5py.File(self.file_path, \"r\") as f:\n                self.trial_keys = sorted(list(f.keys()))\n        except FileNotFoundError:\n            # Handle cases where a subfolder might be missing a split\n            print(f\"Warning: File not found {self.file_path}, creating empty dataset.\")\n            self.trial_keys = []\n\n    def __len__(self):\n        return len(self.trial_keys)\n\n    def __getitem__(self, idx):\n        if self.file is None:\n            self.file = h5py.File(self.file_path, \"r\")\n            \n        trial_key = self.trial_keys[idx]\n        trial_group = self.file[trial_key]\n        \n        x_data = trial_group[self.input_key][:]\n        x = torch.tensor(x_data, dtype=torch.float32)\n        \n        if self.use_augmentation and not self.is_test:\n            x = temporal_mask(x, mask_percentage=0.2)\n        \n        if self.target_key in trial_group:\n            # Assume targets are a 1D array of integer indices (for CTC)\n            y_data = trial_group[self.target_key][:]\n            y = torch.tensor(y_data, dtype=torch.long)\n        else:\n            # Create an empty long tensor as a placeholder for test/dummy targets\n            y = torch.tensor([], dtype=torch.long)\n        \n        if self.is_test:\n            return x, y, trial_key\n        else:\n            return x, y\n\ndef load_datasets():\n    \"\"\"\n    Scans all subfolders in CFG.DATA_DIR and creates combined\n    train, val, and test datasets from all found files.\n    \"\"\"\n    train_datasets = []\n    val_datasets = []\n    test_datasets = []\n    data_dir = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\n    subfolders = [f.path for f in os.scandir(data_dir) if f.is_dir()]\n    print(f\"Found {len(subfolders)} session folders.\")\n    \n    for subfolder_path in subfolders:\n        session_name = os.path.basename(subfolder_path)\n        \n        train_file = os.path.join(subfolder_path, \"data_train.hdf5\")\n        val_file = os.path.join(subfolder_path, \"data_val.hdf5\")\n        test_file = os.path.join(subfolder_path, \"data_test.hdf5\")\n\n        # --- PASS THE use_augmentation FLAG ---\n        # Only apply augmentations to the training set\n        train_set = BrainDataset(train_file, input_key=\"input_features\", target_key=\"seq_class_ids\", is_test=False, use_augmentation=True)\n        val_set = BrainDataset(val_file, input_key=\"input_features\", target_key=\"seq_class_ids\", is_test=False, use_augmentation=False)\n        test_set = BrainDataset(test_file, input_key=\"input_features\", target_key=\"seq_class_ids\", is_test=True, use_augmentation=False) \n        \n        if len(train_set) > 0:\n            train_datasets.append(train_set)\n        if len(val_set) > 0:\n            val_datasets.append(val_set)\n        if len(test_set) > 0:\n            test_datasets.append(test_set)\n            \n    # Combine all individual datasets into one large dataset\n    full_train_dataset = ConcatDataset(train_datasets)\n    full_val_dataset = ConcatDataset(val_datasets)\n    full_test_dataset = ConcatDataset(test_datasets)\n    \n    return full_train_dataset, full_val_dataset, full_test_dataset\n\nprint(\"Loading Train/Val/Test data from session folders...\")\ntrain_dataset, val_dataset, test_dataset = load_datasets()\nprint(\"=\"*40)\nprint(f\"Total Train samples: {len(train_dataset)}\")\nprint(f\"Total Val samples: {len(val_dataset)}\")\nprint(f\"Total Test samples: {len(test_dataset)}\")\nprint(\"=\"*40)\n\n# Check samples\nsample_x_train, sample_y_train = train_dataset[0]\nprint(f\"Train sample X shape: {sample_x_train.shape}\")\nprint(f\"Train sample Y (indices): {sample_y_train}\")\nprint(f\"Train sample Y shape: {sample_y_train.shape}, dtype: {sample_y_train.dtype}\")\n\nif len(test_dataset) > 0:\n    sample_x_test, sample_y_test, _ = test_dataset[0]\n    print(f\"Test sample X shape: {sample_x_test.shape}\")\n    print(f\"Test sample Y (dummy): {sample_y_test}\")\n    print(f\"Test sample Y shape: {sample_y_test.shape}, dtype: {sample_y_test.dtype}\")\nelse:\n    print(f\"Warning: Competition test file not found at {CFG.COMPETITION_TEST_PATH}\")\n    print(\"This is normal. The file will be present during submission.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:05.849402Z","iopub.execute_input":"2025-12-30T07:27:05.849621Z","iopub.status.idle":"2025-12-30T07:27:12.979633Z","shell.execute_reply.started":"2025-12-30T07:27:05.849606Z","shell.execute_reply":"2025-12-30T07:27:12.978859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def clean_y(y):\n    return y[y != 0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:12.980595Z","iopub.execute_input":"2025-12-30T07:27:12.980869Z","iopub.status.idle":"2025-12-30T07:27:12.984461Z","shell.execute_reply.started":"2025-12-30T07:27:12.980846Z","shell.execute_reply":"2025-12-30T07:27:12.983772Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.utils.rnn as rnn_utils\n\ndef custom_collate(batch):\n    \"\"\"\n    Custom collate function for CTC (Sequence-to-Sequence).\n    Pads both x (inputs) and y (targets) and returns their original lengths.\n    `batch` is a list of tuples: (x, y) or (x, y, key)\n    \"\"\"\n    # Check if it's a test batch (has 3 items: x, y, key)\n    is_test = len(batch[0]) == 3\n\n    if is_test:\n        xs, ys, keys = zip(*batch)\n    else:\n        xs, ys = zip(*batch) # Train/val batch\n        \n    # These are the unpadded lengths, required by CTCLoss\n    x_lengths = torch.tensor([len(x) for x in xs], dtype=torch.long)\n    padded_xs = rnn_utils.pad_sequence(xs, batch_first=True, padding_value=0.0)\n    # y_lengths = torch.tensor([len(y) for y in ys], dtype=torch.long)\n    labels = []\n    y_lengths = []\n    for y in ys:\n        y_clean = y[y != 0]\n    \n        labels.append(y_clean)\n        y_lengths.append(len(y_clean))\n    # 2. Pad the 'y' sequences (targets)\n    # We use padding_value=0. This assumes '0' is your 'blank' token index.\n    # This will also correctly handle the empty 'y' tensors from the test set.\n    padded_ys = torch.cat(labels)\n    y_lengths = torch.tensor(y_lengths)\n    \n    if is_test:\n        return padded_xs, padded_ys, x_lengths, y_lengths, keys\n    else:\n        return padded_xs, padded_ys, x_lengths, y_lengths\n\n# Create the DataLoaders using this new collate function\n\n# filtered_indices = [\n#     i for i in range(len(train_dataset))\n#     if len(clean_y(train_dataset[i][1])) <= 15\n# ]\n\n# subset = torch.utils.data.Subset(train_dataset, filtered_indices)\n\ntrain_loader = DataLoader(\n    train_dataset, \n    batch_size=64, \n    shuffle=True, \n    collate_fn=custom_collate,\n)\nval_loader = DataLoader(\n    val_dataset, \n    batch_size=256, \n    shuffle=False, \n    collate_fn=custom_collate \n)\ntest_loader = DataLoader(\n    test_dataset, \n    batch_size=32, \n    shuffle=False, \n    collate_fn=custom_collate\n)\n\nprint(\"DataLoaders created with CTC-ready padding.\")\n\n# --- Let's check a sample batch ---\ntry:\n    x_batch, y_batch, x_len, y_len = next(iter(train_loader))\n    print(\"\\nChecking one batch from train_loader:\")\n    print(f\"  x_batch shape: {x_batch.shape}\")\n    print(f\"  y_batch shape: {y_batch.shape}\")\n    print(f\"  x_lengths shape: {x_len.shape}, sample: {x_len[:5]}\")\n    print(f\"  y_lengths shape: {y_len.shape}, sample: {y_len[:5]}\")\nexcept Exception as e:\n    print(f\"\\nCould not get batch from train_loader (is it empty?): {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:12.985139Z","iopub.execute_input":"2025-12-30T07:27:12.985367Z","iopub.status.idle":"2025-12-30T07:27:15.141010Z","shell.execute_reply.started":"2025-12-30T07:27:12.985347Z","shell.execute_reply":"2025-12-30T07:27:15.140245Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import clear_output\nimport matplotlib.pyplot as plt\n%matplotlib inline","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.142026Z","iopub.execute_input":"2025-12-30T07:27:15.142263Z","iopub.status.idle":"2025-12-30T07:27:15.147148Z","shell.execute_reply.started":"2025-12-30T07:27:15.142246Z","shell.execute_reply":"2025-12-30T07:27:15.146595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def ctc_greedy_decode(batch_preds, blank=0):\n    decoded = []\n    for pred in batch_preds:\n        prev = blank\n        seq = []\n        for p in pred:\n            p = p.item()\n            if p != prev and p != blank:\n                seq.append(p)\n            prev = p\n        decoded.append(seq)\n    return decoded","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.147839Z","iopub.execute_input":"2025-12-30T07:27:15.148076Z","iopub.status.idle":"2025-12-30T07:27:15.158866Z","shell.execute_reply.started":"2025-12-30T07:27:15.148053Z","shell.execute_reply":"2025-12-30T07:27:15.158150Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def levenshtein(seq1, seq2):\n    m, n = len(seq1), len(seq2)\n    dp = [[0] * (n + 1) for _ in range(m + 1)]\n\n    for i in range(m + 1):\n        dp[i][0] = i\n    for j in range(n + 1):\n        dp[0][j] = j\n\n    for i in range(1, m + 1):\n        for j in range(1, n + 1):\n            cost = 0 if seq1[i - 1] == seq2[j - 1] else 1\n            dp[i][j] = min(\n                dp[i - 1][j] + 1,\n                dp[i][j - 1] + 1,\n                dp[i - 1][j - 1] + cost\n            )\n    return dp[m][n]\n\n\ndef compute_batch_cer(decoded_preds, true_labels):\n    total_edits = 0\n    total_chars = 0\n\n    for pred, true in zip(decoded_preds, true_labels):\n        total_edits += levenshtein(pred, true)\n        total_chars += len(true)\n\n    return total_edits, total_chars","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.159631Z","iopub.execute_input":"2025-12-30T07:27:15.159863Z","iopub.status.idle":"2025-12-30T07:27:15.173628Z","shell.execute_reply.started":"2025-12-30T07:27:15.159842Z","shell.execute_reply":"2025-12-30T07:27:15.173017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_ce_targets(log_probs, labels, label_lengths, blank_idx=0):\n    \"\"\"\n    log_probs: (T, B, C)\n    labels: (B, L)\n    returns aligned_targets: (B, T)\n    \"\"\"\n\n    B, T, C = log_probs.shape\n    aligned_targets = torch.full((B, T), blank_idx, dtype=torch.long, device=log_probs.device)\n    # Greedy prediction: (T, B)\n    # p = 0\n    # for b in range(B):\n    #     l = label_lengths[b]\n    #     label = labels[p:l+p] \n    #     p = l\n    #     augmented_probs = log_probs[:,:,label] \n    #     augmented_probs += 0.4\n    #     log_probs[:,:,label] = augmented_probs \n\n    beam_decoded = beam_decoder(log_probs.cpu())\n    pred = []\n    for b in range(B):\n        pred.append(beam_decoded[b][0].tokens.tolist())\n    p = 0\n    for b in range(B):\n        l = label_lengths[b]\n        label = labels[p:l+p] \n        p = l\n\n        pred_seq = pred[b]\n        # collapse consecutive duplicates and blanks in pred\n        prev = None\n        collapsed = []\n        for p in pred_seq:\n            if p != blank_idx and p != prev:\n                collapsed.append(p)\n            prev = p\n        \n        # collapsed is our approximated predicted sequence\n        # we want to align collapsed to label\n\n        i = 0  # pointer in collapsed\n        j = 0  # pointer in label\n        \n        # -------- forward pass (earliest placement) --------\n        for t in range(T):\n            if i < len(collapsed) and j < len(label) and collapsed[i] == label[j].item():\n                aligned_targets[b, t] = label[j].item()\n                i += 1\n                j += 1\n            else:\n                aligned_targets[b, t] = blank_idx\n        \n        \n        # -------- reset pointers for backward pass --------\n        i = len(collapsed) - 1\n        j = len(label) - 1\n        \n        \n        # -------- backward pass (latest placement check) --------\n        for t in range(T - 1, -1, -1):\n            if i >= 0 and j >= 0 and collapsed[i] == label[j].item():\n                # keep only if forward already placed same label here\n                if aligned_targets[b, t] == blank_idx:\n                    aligned_targets[b, t] = label[j].item()\n                i -= 1\n                j -= 1\n            else:\n                aligned_targets[b, t] = blank_idx\n\n\n    return aligned_targets, pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.174272Z","iopub.execute_input":"2025-12-30T07:27:15.174495Z","iopub.status.idle":"2025-12-30T07:27:15.187383Z","shell.execute_reply.started":"2025-12-30T07:27:15.174475Z","shell.execute_reply":"2025-12-30T07:27:15.186708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def uniform_alignment(labels, T, blank=0):\n    aligned = torch.full((T,), blank, dtype=torch.long)\n    L = len(labels)\n    if L == 0:\n        return aligned\n    positions = torch.linspace(0, T - 1, steps=L).long()\n    aligned[positions] = labels\n    return aligned","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.187979Z","iopub.execute_input":"2025-12-30T07:27:15.188302Z","iopub.status.idle":"2025-12-30T07:27:15.203363Z","shell.execute_reply.started":"2025-12-30T07:27:15.188284Z","shell.execute_reply":"2025-12-30T07:27:15.202681Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate(model, ctc_fn, ce_fn, dataloader):\n    correct, total = 0, 0\n    model.eval()\n    chunk_length = 1323000  # 60 seconds * 22050 Hz\n    predicted_all = []\n    total_loss = 0\n    epoch_char_errors = 0\n    epoch_target_chars = 0\n    epoch_loss = 0.0\n    with torch.no_grad():\n        # for _ in range(9657):\n        for i, (inputs, labels, x_lengths, y_lengths) in enumerate(dataloader):\n            inputs = inputs.to(device1)\n            labels = labels.to(device0)\n            x_lengths = x_lengths.to(device0) // 16\n            y_lengths = y_lengths.to(device0)\n            outputs = model(inputs)\n            log_probs = torch.log_softmax(outputs, dim=2)\n            T = log_probs.size(0)\n            ctc_loss = ctc_fn(\n                log_probs.permute(1,0,2),\n                labels,\n                x_lengths,\n                y_lengths\n            )\n\n            # predicted = torch.argmax(log_probs, dim=-1).permute(1, 0)  # [B, T]\n            ctc_hypothesis = beam_decoder(log_probs.cpu())\n            decoded_preds = []\n            for b in range(64):\n                decoded_preds.append(ctc_hypothesis[b][0].tokens.tolist())\n            # ce_loss = ce_fn(outputs.permute(0, 2, 1), predicted) \n            # loss = ctc_loss + (ce_loss * alpha)\n            loss = ctc_loss\n\n            p = 0\n            true_labels = []\n            for l in y_lengths:\n                true_labels.append(labels[p:p+l].tolist())\n                p = l\n\n            batch_edits, batch_chars = compute_batch_cer(\n                decoded_preds, true_labels\n            )\n\n            epoch_char_errors += batch_edits\n            epoch_target_chars += batch_chars\n            \n            batch_correct = 0\n            batch_total = 0\n            \n            for pred_seq, true_seq in zip(decoded_preds, true_labels):\n                L = min(len(pred_seq), len(true_seq))\n                for k in range(L):\n                    if pred_seq[k] == true_seq[k]:\n                        batch_correct += 1\n                batch_total += len(true_seq)\n\n            lengths = [len(seq) for seq in decoded_preds]\n            avg_length = sum(lengths) / len(decoded_preds)\n            \n            del inputs, labels, log_probs, decoded_preds\n            correct += batch_correct\n            total_loss += loss\n            total += batch_total\n         \n        # print(i+1)\n            # if (i+1) % 1000 == 0:\n                # print(100 * correct / total)\n    avg_loss = total_loss / len(val_loader)\n    epoch_cer = epoch_char_errors / max(1, epoch_target_chars)\n\n    model.train()\n    return epoch_cer, avg_loss.cpu(), avg_length","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.204035Z","iopub.execute_input":"2025-12-30T07:27:15.204255Z","iopub.status.idle":"2025-12-30T07:27:15.216574Z","shell.execute_reply.started":"2025-12-30T07:27:15.204231Z","shell.execute_reply":"2025-12-30T07:27:15.215859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"phoneme_counts = {17: 10561, 31: 13049, 29: 7054, 40: 50645, 23: 9368, 1: 3403, \n                 2: 4373, 4: 2341, 21: 6429, 11: 4589, 18: 6278, 32: 1003,\n                 24: 2073, 36: 4100, 12: 3140, 20: 5178, 3: 12662, 7: 2593,\n                 5: 1316, 0: 3821756, 10: 5258, 13: 2699, 27: 2939, 28: 6865,\n                  22: 5125, 8: 736, 15: 2110, 35: 3334, 9: 6260, 16: 2942, 38: 4348,\n             25: 2663,\n             6: 5337,\n             14: 2645,\n             30: 736,\n             34: 4801,\n             33: 1016,\n             19: 831,\n             39: 436,\n             37: 2407,\n             26: 601}\n# [3, 1, 38, 16, 21, 16, 7, 1, 16]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.217353Z","iopub.execute_input":"2025-12-30T07:27:15.217660Z","iopub.status.idle":"2025-12-30T07:27:15.233505Z","shell.execute_reply.started":"2025-12-30T07:27:15.217643Z","shell.execute_reply":"2025-12-30T07:27:15.232811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"counts = np.zeros(len(LOGIT_TO_PHONEME))\nfor cls, cnt in phoneme_counts.items():\n    counts[cls] = cnt\n\nweights = (counts.sum() / (len(LOGIT_TO_PHONEME) * counts)) ** 0.3\nweights = torch.tensor(weights, device = device1, dtype=torch.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.234099Z","iopub.execute_input":"2025-12-30T07:27:15.234402Z","iopub.status.idle":"2025-12-30T07:27:15.381634Z","shell.execute_reply.started":"2025-12-30T07:27:15.234382Z","shell.execute_reply":"2025-12-30T07:27:15.381044Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.384828Z","iopub.execute_input":"2025-12-30T07:27:15.385050Z","iopub.status.idle":"2025-12-30T07:27:15.645287Z","shell.execute_reply.started":"2025-12-30T07:27:15.385032Z","shell.execute_reply":"2025-12-30T07:27:15.644333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install flashlight-text","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:15.645975Z","iopub.execute_input":"2025-12-30T07:27:15.646359Z","iopub.status.idle":"2025-12-30T07:27:20.332810Z","shell.execute_reply.started":"2025-12-30T07:27:15.646334Z","shell.execute_reply":"2025-12-30T07:27:20.332089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchaudio.models.decoder import ctc_decoder\n\nbeam_decoder = ctc_decoder(\n    lexicon=None,\n    tokens=LOGIT_TO_PHONEME,\n    beam_size=15,\n    blank_token=\"BLANK\",\n    sil_token = ' | '\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:20.333805Z","iopub.execute_input":"2025-12-30T07:27:20.334080Z","iopub.status.idle":"2025-12-30T07:27:21.010261Z","shell.execute_reply.started":"2025-12-30T07:27:20.334042Z","shell.execute_reply":"2025-12-30T07:27:21.009690Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_loop(model, lr, ctc_fn, ce_fn, batch_size, alpha=0.1):\n    # Ensure all models are on their assigned device outside this function!\n    # model on device0, model_2 experts and Feature Extractor on device1\n\n    optim = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=3e-6)\n\n    \n    expert_keys_used = []\n    losses = [0]\n    val_losses = [0]\n    rare_losses = [0]\n    accuracies = [0]\n    val_accuracies = [0]\n    avg_lengths = [0]\n    val_accuracy, val_loss = 0, 0\n    use_ce = True\n    blank_weight = 4.0\n    for epoch in range(2):\n        torch.cuda.empty_cache()\n        print(\"Current Epoch:\", epoch+1)\n        running_loss = 0.0\n        correct, total = 0, 0\n        optim.zero_grad()\n        expert_keys_used = []\n        epoch_char_errors = 0\n        epoch_target_chars = 0\n        epoch_loss = 0.0\n        batch_cer = 0\n        batch_length, real_batch_length = 0, 0\n        for i, (inputs, labels, x_lengths, y_lengths) in enumerate(train_loader):\n            # assert (x_lengths >= y_lengths).all()\n            # if i+1 >= 5:\n            #     use_ce=True\n            inputs = inputs.to(device1)\n            labels = labels.to(device0)\n            x_lengths = ((x_lengths.to(device0) - 1) // 16) + 1\n            y_lengths = y_lengths.to(device0)\n            outputs = model(inputs)\n            logits = outputs\n            logits[:,:,40] -= blank_weight\n            logits[:,:,0] -= blank_weight\n            # FUCKS Up the logs and collapses the average length\n            # logits = logits + scaled_w.unsqueeze(0).unsqueeze(0).to(device0)\n            log_probs = torch.log_softmax(logits, dim=2)\n            # log_probs = torch.clamp(log_probs, min=-20, max=0)\n            # print(x_lengths)\n            # ctc_loss = ctc_fn(\n            #     log_probs.permute(1,0,2),\n            #     labels,\n            #     x_lengths,\n            #     y_lengths\n            # )\n\n            # predicted = torch.argmax(, dim=-1).permute(1, 0)  # [B, T]\n            ctc_hypothesis = beam_decoder(log_probs.cpu())\n            decoded_preds = []\n            B, _, _ = inputs.shape\n\n            for b in range(B):\n                decoded_preds.append(ctc_hypothesis[b][0].tokens.tolist())\n            with torch.no_grad():\n                aligned, pred = get_ce_targets(log_probs.clone().detach(), labels, y_lengths)\n            \n            lengths = [len(seq) for seq in decoded_preds]\n            batch_length += sum(lengths) / len(decoded_preds)\n            real_batch_length += sum(y_lengths) / len(y_lengths)\n            \n            # del_penalty = F.relu(real_batch_length - batch_length).mean()\n            if aligned.sum()!=0:\n                    # print(outputs.shape)\n                    ce_loss = ce_fn(\n                        outputs.permute(0, 2, 1),\n                        aligned\n                    )\n            # if use_ce:\n            #     # print(ce_loss)\n            \n            #     print(f\"Step {i+1} CE Loss : {ce_loss}\")\n            #     loss = ctc_loss + alpha * ce_loss\n            # else:\n            #     # loss = ctc_loss\n            #     ctc_loss = (ctc_loss + 0.15 * dup_penalty) \n            \n\n            loss = ce_loss / batch_size\n\n            p = 0\n            true_labels = []\n            for l in y_lengths:\n                true_labels.append(labels[p:p+l].tolist())\n                p = l\n\n            batch_edits, batch_chars = compute_batch_cer(\n                decoded_preds, true_labels\n            )\n\n            epoch_char_errors += batch_edits\n            epoch_target_chars += batch_chars\n            \n            batch_cer += batch_edits / max(1, batch_chars)\n            del inputs, labels, log_probs\n            \n            running_loss += loss\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n            if (i+1)%batch_size == 0:\n                optim.step()\n                optim.zero_grad()\n                torch.cuda.empty_cache()\n                # --- Model 1 Tracking ---\n                avg_length = batch_length / batch_size\n                real_avg_length = real_batch_length / batch_size\n                avg_lengths.append(avg_length)\n                \n                batch_length, real_batch_length = 0, 0\n                \n                avg_loss = running_loss \n                running_loss = 0.0\n                acc = batch_cer / batch_size\n                batch_cer = 0\n                correct, total = 0, 0\n                losses.append(avg_loss.to(\"cpu\").detach().numpy())\n                accuracies.append(acc)\n    \n                print(f\"Step {(i+1)//batch_size} Done | CER : {acc} | Loss : {avg_loss} | Avg Length : {avg_length} | Real AVG Length : {real_avg_length}\", flush=True) # Added flush=True for immediate log\n                if ((i+1)//batch_size)%10 == 0:\n                    val_accuracy, val_loss, val_avg_length = evaluate(model, ctc_fn, ce_fn, val_loader)\n                    print(f\"Val CER: {val_accuracy} | Avg Length: {val_avg_length} | Loss: {val_loss}\", flush=True)\n                # if ((i+1)//batch_size) == 20:\n                #     return\n                val_losses.append(val_loss)\n                val_accuracies.append(val_accuracy)\n                    \n                clear_output(wait=True)\n                plt.clf()\n    \n                fig, ax = plt.subplots(3, 1, figsize=(14, 8))\n                ax[0].plot(losses, label='Training Loss')\n                ax[0].plot(val_losses, label='Val Loss')\n                ax[0].set_title(f\"Epoch {epoch+1}\")\n                ax[0].set_xlabel(f\"{32}-Batch Steps\")\n                ax[0].set_ylabel(\"Loss Value\")\n                ax[0].legend()\n                ax[0].grid(True)\n    \n                ax[1].plot(accuracies, label='Train CER')\n                ax[1].plot(val_accuracies, label='Val CER')\n                ax[1].set_title(f\"Epoch {epoch+1}\")\n                ax[1].set_xlabel(f\"{32}-Batch Steps\")\n                ax[1].set_ylabel(\"Accuracy\")\n                ax[1].legend()\n                ax[1].grid(True)\n    \n                # print(avg_lengths)\n                ax[2].plot(avg_lengths, label='AVG Length')\n                ax[2].set_title(f\"Epoch {epoch+1}\")\n                ax[2].set_xlabel(f\"{32}-Batch Steps\")\n                ax[2].set_ylabel(\"Length\")\n                ax[2].legend()\n                ax[2].grid(True)\n                \n                \n                plt.tight_layout()\n                plt.show()\n                plt.pause(0.001)\n            \n        epoch_cer = epoch_char_errors / max(1, epoch_target_chars)\n\n        print(f\"Epoch {epoch+1} | CER: {epoch_cer:.4f}\")\n\n    print('Finished Training')\n    plt.ioff()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:33:24.541524Z","iopub.execute_input":"2025-12-30T07:33:24.541819Z","iopub.status.idle":"2025-12-30T07:33:24.562004Z","shell.execute_reply.started":"2025-12-30T07:33:24.541798Z","shell.execute_reply":"2025-12-30T07:33:24.561057Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# elm = ELM(\n#     num_input = 512,\n#     num_output = len(LOGIT_TO_PHONEME),\n#     num_memory=512,\n#     mlp_num_layers=5,\n#     memory_tau_min=1,\n#     memory_tau_max=784,\n#     num_branch=8,\n#     num_synapse_per_branch=64,\n#     input_to_synapse_routing = 'neuronio_routing'\n# ).to(device0)\n# # elm = ELM(**readout_config, num_input=512).to(device1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:21.025593Z","iopub.execute_input":"2025-12-30T07:27:21.025932Z","iopub.status.idle":"2025-12-30T07:27:21.041382Z","shell.execute_reply.started":"2025-12-30T07:27:21.025910Z","shell.execute_reply":"2025-12-30T07:27:21.040753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"nodes_config = {\n    \"input_to_synapse_routing\": \"neuronio_routing\",\n    \"learn_memory_tau\": False,\n    \"memory_tau_max\": 200.0,\n    \"memory_tau_min\": 1.0,\n    \"mlp_activation\": \"silu\",\n    \"num_branch\": 16,\n    \"mlp_num_layers\": 5,\n    \"num_input\": 256,\n    \"num_memory\": 128,\n    \"num_output\": 256,\n    \"num_synapse_per_branch\": 100\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:21.042064Z","iopub.execute_input":"2025-12-30T07:27:21.042320Z","iopub.status.idle":"2025-12-30T07:27:21.053198Z","shell.execute_reply.started":"2025-12-30T07:27:21.042297Z","shell.execute_reply":"2025-12-30T07:27:21.052494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"reservoir_params = {\n    'internal_nodes_params' : nodes_config.copy(), # Dict for internal nodes param\n    'num_internal_nodes' : 50,\n    'internal_output_dim' : 1, # Number of output per internal nodes\n    'internal_connectivity' : 0.3\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:27:21.054598Z","iopub.execute_input":"2025-12-30T07:27:21.054802Z","iopub.status.idle":"2025-12-30T07:27:21.065579Z","shell.execute_reply.started":"2025-12-30T07:27:21.054787Z","shell.execute_reply":"2025-12-30T07:27:21.064892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CResT(\n    num_input=512,\n    num_output=len(LOGIT_TO_PHONEME),\n    reservoir_params=nodes_config.copy(),\n    readout_mlp_layers=3,\n    conv_feature_out=512,\n).to(device1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:28:25.559712Z","iopub.execute_input":"2025-12-30T07:28:25.560354Z","iopub.status.idle":"2025-12-30T07:28:25.594000Z","shell.execute_reply.started":"2025-12-30T07:28:25.560325Z","shell.execute_reply":"2025-12-30T07:28:25.593397Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATH = \"/kaggle/input/brain-to-speech-using-elm/model_weights.pth\"\ncheckpoint = torch.load(PATH, weights_only=False)\nmodel.load_model(checkpoint)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:28:26.355791Z","iopub.execute_input":"2025-12-30T07:28:26.356314Z","iopub.status.idle":"2025-12-30T07:28:26.376800Z","shell.execute_reply.started":"2025-12-30T07:28:26.356291Z","shell.execute_reply":"2025-12-30T07:28:26.376225Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# for param in model.reservoir.parameters():\n#     param.requires_grad=False\n# for param in model.conv.parameters():\n#     param.requires_grad=False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:28:30.327650Z","iopub.execute_input":"2025-12-30T07:28:30.328184Z","iopub.status.idle":"2025-12-30T07:28:30.332041Z","shell.execute_reply.started":"2025-12-30T07:28:30.328162Z","shell.execute_reply":"2025-12-30T07:28:30.331352Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchinfo import summary","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:28:51.323562Z","iopub.execute_input":"2025-12-30T07:28:51.323852Z","iopub.status.idle":"2025-12-30T07:28:51.343238Z","shell.execute_reply.started":"2025-12-30T07:28:51.323832Z","shell.execute_reply":"2025-12-30T07:28:51.342677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(summary(model))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:28:52.190356Z","iopub.execute_input":"2025-12-30T07:28:52.190941Z","iopub.status.idle":"2025-12-30T07:28:52.196945Z","shell.execute_reply.started":"2025-12-30T07:28:52.190906Z","shell.execute_reply":"2025-12-30T07:28:52.196248Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lr = 1e-5\nctc_fn = nn.CTCLoss(blank=0, zero_infinity=True, reduction='mean').to(device0)\nce_fn = nn.CrossEntropyLoss(ignore_index=0, weight=weights).to(device0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:29:01.025692Z","iopub.execute_input":"2025-12-30T07:29:01.026442Z","iopub.status.idle":"2025-12-30T07:29:01.030653Z","shell.execute_reply.started":"2025-12-30T07:29:01.026417Z","shell.execute_reply":"2025-12-30T07:29:01.029938Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:29:01.484368Z","iopub.execute_input":"2025-12-30T07:29:01.484947Z","iopub.status.idle":"2025-12-30T07:29:01.488419Z","shell.execute_reply.started":"2025-12-30T07:29:01.484927Z","shell.execute_reply":"2025-12-30T07:29:01.487633Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loop(model, lr, ctc_fn, ce_fn, batch_size=1)\n# evaluate(model, ctc_fn, ce_fn, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:33:33.778210Z","iopub.execute_input":"2025-12-30T07:33:33.778723Z","iopub.status.idle":"2025-12-30T07:34:29.410294Z","shell.execute_reply.started":"2025-12-30T07:33:33.778702Z","shell.execute_reply":"2025-12-30T07:34:29.409647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save this motherfucker before doing anything else to \n# prevent ruining the whole run like before\nPATH = \"model_weights.pth\"\ntorch.save(model.save_model(), PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T13:02:49.318256Z","iopub.execute_input":"2025-12-28T13:02:49.319023Z","iopub.status.idle":"2025-12-28T13:02:49.339243Z","shell.execute_reply.started":"2025-12-28T13:02:49.318988Z","shell.execute_reply":"2025-12-28T13:02:49.338256Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"(inputs, labels, x_lengths, y_lengths) = next(iter(train_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:09.569814Z","iopub.execute_input":"2025-12-30T07:32:09.570312Z","iopub.status.idle":"2025-12-30T07:32:10.847956Z","shell.execute_reply.started":"2025-12-30T07:32:09.570283Z","shell.execute_reply":"2025-12-30T07:32:10.847334Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"logits = model(inputs.to(device1))\nlog_probs = torch.softmax(logits, dim=-1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:10.849250Z","iopub.execute_input":"2025-12-30T07:32:10.849592Z","iopub.status.idle":"2025-12-30T07:32:11.097171Z","shell.execute_reply.started":"2025-12-30T07:32:10.849568Z","shell.execute_reply":"2025-12-30T07:32:11.096564Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"log_probs.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:12.963206Z","iopub.execute_input":"2025-12-30T07:32:12.963952Z","iopub.status.idle":"2025-12-30T07:32:12.969097Z","shell.execute_reply.started":"2025-12-30T07:32:12.963925Z","shell.execute_reply":"2025-12-30T07:32:12.968387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"beam_decoded = beam_decoder(log_probs.cpu())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:13.485685Z","iopub.execute_input":"2025-12-30T07:32:13.486421Z","iopub.status.idle":"2025-12-30T07:32:14.153942Z","shell.execute_reply.started":"2025-12-30T07:32:13.486395Z","shell.execute_reply":"2025-12-30T07:32:14.153139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"true_labels = []\np = 0 \nfor l in y_lengths:\n    true_labels.append(labels[p:p+l])\n    p = l","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:14.154957Z","iopub.execute_input":"2025-12-30T07:32:14.155262Z","iopub.status.idle":"2025-12-30T07:32:14.160133Z","shell.execute_reply.started":"2025-12-30T07:32:14.155239Z","shell.execute_reply":"2025-12-30T07:32:14.159560Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"decoded_preds = []\nfor b in range(64):\n    decoded_preds.append(beam_decoded[b][0].tokens.tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:16.302494Z","iopub.execute_input":"2025-12-30T07:32:16.303244Z","iopub.status.idle":"2025-12-30T07:32:16.306882Z","shell.execute_reply.started":"2025-12-30T07:32:16.303210Z","shell.execute_reply":"2025-12-30T07:32:16.306166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i, (d, t) in enumerate(zip(decoded_preds, true_labels)):\n    print(f\"Predicted: {d}\")\n    print(f\"True: {t}\")\n    # break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:28.170546Z","iopub.execute_input":"2025-12-30T07:32:28.170834Z","iopub.status.idle":"2025-12-30T07:32:28.188288Z","shell.execute_reply.started":"2025-12-30T07:32:28.170812Z","shell.execute_reply":"2025-12-30T07:32:28.187538Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"aligned, pred = get_ce_targets(log_probs, labels, y_lengths)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:17.565658Z","iopub.execute_input":"2025-12-30T07:32:17.566366Z","iopub.status.idle":"2025-12-30T07:32:18.615252Z","shell.execute_reply.started":"2025-12-30T07:32:17.566343Z","shell.execute_reply":"2025-12-30T07:32:18.614631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"aligned[aligned != 0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T07:32:18.617680Z","iopub.execute_input":"2025-12-30T07:32:18.617895Z","iopub.status.idle":"2025-12-30T07:32:18.623524Z","shell.execute_reply.started":"2025-12-30T07:32:18.617879Z","shell.execute_reply":"2025-12-30T07:32:18.622863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"un = set()\nfor p in aligned:\n    print(p)\n    break\n    for u in p.unique():\n        un.add(u.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T05:22:24.199999Z","iopub.execute_input":"2025-12-29T05:22:24.200677Z","iopub.status.idle":"2025-12-29T05:22:24.207928Z","shell.execute_reply.started":"2025-12-29T05:22:24.200650Z","shell.execute_reply":"2025-12-29T05:22:24.207142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"un","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T05:22:26.733877Z","iopub.execute_input":"2025-12-29T05:22:26.734566Z","iopub.status.idle":"2025-12-29T05:22:26.738984Z","shell.execute_reply.started":"2025-12-29T05:22:26.734528Z","shell.execute_reply":"2025-12-29T05:22:26.738366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lengths = [len(seq) for seq in decoded_preds]\navg_length = sum(lengths) / len(decoded_preds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T05:22:27.179146Z","iopub.execute_input":"2025-12-29T05:22:27.179432Z","iopub.status.idle":"2025-12-29T05:22:27.183344Z","shell.execute_reply.started":"2025-12-29T05:22:27.179409Z","shell.execute_reply":"2025-12-29T05:22:27.182763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"avg_length","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T05:22:28.919253Z","iopub.execute_input":"2025-12-29T05:22:28.919539Z","iopub.status.idle":"2025-12-29T05:22:28.924471Z","shell.execute_reply.started":"2025-12-29T05:22:28.919517Z","shell.execute_reply":"2025-12-29T05:22:28.923745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test_evaluate(model, dataloader, phoneme=LOGIT_TO_PHONEME):\n    correct, total = 0, 0\n    model.eval()\n    chunk_length = 1323000  # 60 seconds * 22050 Hz\n    predicted_all = []\n    with torch.no_grad():\n        # for _ in range(9657):\n        for i, (inputs, labels, x_lengths, y_lengths, keys) in enumerate(dataloader):\n            outputs = model(inputs.to(\"cuda\"))\n            prob = torch.softmax(outputs, dim=-1)\n            predicted = torch.argmax(prob, dim=-1)\n            \n            for i, row in enumerate(predicted):\n                sequence = []\n                for phe in row:\n                    if phe < len(phoneme):\n                        sequence.append(phoneme[phe])\n                    else:\n                        sequence.append(phoneme[-1]) # Incase model outputs out of bound indices \n                predicted_all.append(sequence)\n         \n    model.train()\n    return predicted_all","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-28T13:03:47.247681Z","iopub.execute_input":"2025-12-28T13:03:47.248420Z","iopub.status.idle":"2025-12-28T13:03:47.253727Z","shell.execute_reply.started":"2025-12-28T13:03:47.248394Z","shell.execute_reply":"2025-12-28T13:03:47.252947Z"}},"outputs":[],"execution_count":null}]}