{
  "id": 665582,
  "title": "data augmentation code",
  "url": "/competitions/brain-to-text-25/discussion/665582",
  "author_name": "Cyrus",
  "post_date": "2026-01-02T10:35:20.532000",
  "votes": 0,
  "comment_count": 0,
  "views": 0,
  "content": "<p>I did not invest much time in this work, particularly in selecting the transformer encoder–decoder, which underperformed compared to RNNs, especially when tokenizing text rather than phonemes. However, I believe I did a solid job on data augmentation, which may be useful for other researchers. here is the code: <a href=\"https://www.kaggle.com/code/jamalsaeedi/b2t-infer\" target=\"_blank\">https://www.kaggle.com/code/jamalsaeedi/b2t-infer</a> </p>\n<p>`</p>\n<p>class PhysioAwareAugment(nn.Module):</p>\n<pre><code>def __init__(self, num_electrodes: int = 256, bin_duration_ms: int = 20):\n    super().__init__()\n    self.num_electrodes = num_electrodes\n    self.bin_duration_ms = bin_duration_ms\n\n    # Array definitions\n    self.arrays = {\n        'ventral_6v': (0, 64),\n        'area_4': (64, 128),\n        'area_55b': (128, 192),\n        'dorsal_6v': (192, 256),\n    }\n\n    self.critical_arrays = ['area_4', 'area_55b']\n\n    # Feature type indices\n    self.TC = slice(0, 256)\n    self.SBP = slice(256, 512)\n\n    # Gaussian smoothing parameters\n    self.smooth_kernel_std = 2.0\n    self.smooth_kernel_size = 100\n    self.has_scipy = HAS_SCIPY\n\ndef _ensure_batch_dim(self, x: torch.Tensor) -&gt; Tuple[torch.Tensor, bool]:\n    \"\"\"Ensure input has batch dimension\"\"\"\n    if x.dim() == 2:  # (T, C)\n        return x.unsqueeze(0), True  # (1, T, C)\n    elif x.dim() == 3:  # (B, T, C)\n        return x, False\n    else:\n        raise ValueError(f\"Expected 2D or 3D input, got shape {x.shape}\")\n\ndef _restore_shape(self, x: torch.Tensor, squeeze: bool) -&gt; torch.Tensor:\n    \"\"\"Restore original shape after augmentation\"\"\"\n    return x.squeeze(0) if squeeze else x\n\ndef _temporal_warp(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"Clean, efficient temporal warping that actually works\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    if T &lt; 10:  # Safety check for short sequences\n        return x\n\n    # Create warp field with 5 control points\n    control_pts = torch.linspace(0, 1, 5, device=device)\n    offsets = (torch.rand(5, device=device) - 0.5) * 0.15\n    warped_pts = torch.clamp(control_pts + offsets, 0, 1)\n    warped_pts, _ = torch.sort(warped_pts)  # Ensure monotonic\n\n    # Interpolate to get full warp field (shape: [T])\n    warp_field = F.interpolate(\n        warped_pts.view(1, 1, -1),\n        size=T,\n        mode='linear',\n        align_corners=True\n    ).squeeze(0).squeeze(0)  # Ensure shape is [T]\n\n    # Map to indices and clamp\n    indices = warp_field * (T - 1)\n    indices = torch.clamp(indices, 0, T - 1)\n\n    # Get floor and ceiling indices\n    floor_idx = torch.floor(indices).long()\n    ceil_idx = torch.ceil(indices).long()\n    ceil_idx = torch.clamp(ceil_idx, 0, T - 1)\n\n    # Interpolation weights (shape: [T])\n    weight = indices - floor_idx.float()\n\n    # Create output tensor\n    x_warped = torch.zeros_like(x)\n\n    # Vectorized interpolation per batch\n    for b in range(B):\n        # For each batch, interpolate all channels at once\n        floor_vals = x[b, floor_idx, :]  # [T, C]\n        ceil_vals = x[b, ceil_idx, :]    # [T, C]\n\n        # Expand weight to match channel dimension: [T] -&gt; [T, 1]\n        weight_expanded = weight.view(-1, 1)\n\n        # Interpolate: (1-weight)*floor + weight*ceil\n        x_warped[b] = floor_vals * \\\n            (1 - weight_expanded) + ceil_vals * weight_expanded\n\n    return x_warped\n\ndef _gauss_smooth(self, inputs: torch.Tensor, smooth_kernel_std: float = None,\n                  smooth_kernel_size: int = None, padding: str = 'same') -&gt; torch.Tensor:\n    \"\"\"\n    Applies 1D Gaussian smoothing with proper fallback if scipy is not available.\n\n    Args:\n        inputs (tensor): B x T x N tensor\n        smooth_kernel_std (float): Standard deviation of Gaussian kernel\n        smooth_kernel_size (int): Size of Gaussian kernel\n        padding (str): Padding mode ('same' or 'valid')\n    \"\"\"\n    if smooth_kernel_std is None:\n        smooth_kernel_std = self.smooth_kernel_std\n    if smooth_kernel_size is None:\n        smooth_kernel_size = self.smooth_kernel_size\n\n    device = inputs.device\n    B, T, C = inputs.shape\n\n    if self.has_scipy:\n        return self._gauss_smooth_scipy(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n    else:\n        return self._gauss_smooth_torch(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n\ndef _gauss_smooth_scipy(self, inputs: torch.Tensor, device: torch.device,\n                        smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                        padding: str = 'same') -&gt; torch.Tensor:\n    \"\"\"Gaussian smoothing using scipy (more accurate)\"\"\"\n    # Get Gaussian kernel\n    inp = np.zeros(smooth_kernel_size, dtype=np.float32)\n    inp[smooth_kernel_size // 2] = 1\n    gaussKernel = gaussian_filter1d(inp, smooth_kernel_std)\n    validIdx = np.argwhere(gaussKernel &gt; 0.01)\n    gaussKernel = gaussKernel[validIdx]\n    gaussKernel = np.squeeze(gaussKernel / np.sum(gaussKernel))\n\n    # Convert to tensor\n    gaussKernel = torch.tensor(\n        gaussKernel, dtype=torch.float32, device=device)\n    gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n    # Prepare convolution\n    B, T, C = inputs.shape\n    inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n    gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n    # Perform convolution\n    smoothed = F.conv1d(inputs_perm, gaussKernel,\n                        padding=padding, groups=C)\n    return smoothed.permute(0, 2, 1)  # [B, T, C]\n\ndef _gauss_smooth_torch(self, inputs: torch.Tensor, device: torch.device,\n                        smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                        padding: str = 'same') -&gt; torch.Tensor:\n    \"\"\"Gaussian smoothing using pure PyTorch (fallback)\"\"\"\n    # Create Gaussian kernel using PyTorch\n    x = torch.arange(-(smooth_kernel_size//2), smooth_kernel_size //\n                     2 + 1, device=device, dtype=torch.float32)\n    gaussKernel = torch.exp(-x**2 / (2 * smooth_kernel_std**2))\n    gaussKernel = gaussKernel / gaussKernel.sum()\n\n    # Ensure kernel has odd size\n    if gaussKernel.size(0) % 2 == 0:\n        gaussKernel = gaussKernel[:-1]\n\n    gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n    # Prepare convolution\n    B, T, C = inputs.shape\n    inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n    gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n    # Handle padding\n    if padding == 'same':\n        padding_size = gaussKernel.size(2) // 2\n    else:\n        padding_size = 0\n\n    # Perform convolution\n    smoothed = F.conv1d(inputs_perm, gaussKernel,\n                        padding=padding_size, groups=C)\n    return smoothed.permute(0, 2, 1)  # [B, T, C]\n\ndef _create_coupled_electrode_mask(self, x: torch.Tensor, keep_prob: float = 0.85) -&gt; torch.Tensor:\n    \"\"\"Create mask where both TC and SBP features for same electrode are dropped together\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    # Base mask for 256 electrodes\n    elec_mask = torch.bernoulli(torch.ones(\n        self.num_electrodes, device=device) * keep_prob)\n\n    # Add spatial correlation within arrays\n    for array_name, (start_idx, end_idx) in self.arrays.items():\n        center_elec = torch.randint(\n            start_idx, end_idx, (1,), device=device).item()\n        distances = torch.arange(\n            start_idx, end_idx, device=device).float() - center_elec\n        spatial_weights = torch.exp(-torch.abs(distances) / 10.0)\n\n        spatial_bias = 0.25 * spatial_weights\n        p = torch.clamp(keep_prob - spatial_bias, 0.05, 1.0)\n\n        array_mask = torch.bernoulli(p)\n        elec_mask[start_idx:end_idx] = array_mask\n\n    # Expand to 512 features: [TC_mask, SBP_mask]\n    full_mask = torch.cat([elec_mask, elec_mask])\n    return full_mask.view(1, 1, -1)  # (1, 1, 512)\n\ndef _create_array_mask(self, x: torch.Tensor, dropout_prob: float = 0.1) -&gt; torch.Tensor:\n    \"\"\"Create mask that drops entire arrays with anatomical awareness\"\"\"\n    B, T, C = x.shape\n    device = x.device\n    mask = torch.ones(512, device=device)\n\n    for array_name, (start_idx, end_idx) in self.arrays.items():\n        is_critical = array_name in self.critical_arrays\n        array_dropout_prob = dropout_prob * 0.3 if is_critical else dropout_prob\n\n        if torch.rand(1, device=device) &lt; array_dropout_prob:\n            mask[start_idx:end_idx] = 0  # TC features\n            mask[start_idx + 256:end_idx + 256] = 0  # SBP features\n\n    return mask.view(1, 1, -1)  # (1, 1, 512)\n\ndef _apply_feature_specific_noise(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"Apply physiologically appropriate noise to different feature types\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    tc_part = x[:, :, self.TC].clone()\n    sbp_part = x[:, :, self.SBP].clone()\n\n    # Threshold Crossings (TC): multiplicative noise\n    if torch.rand(1, device=device) &lt; 0.3:\n        scale_factor = 0.9 + torch.rand(1, device=device).item() * 0.2\n        tc_part = tc_part * scale_factor\n\n        if torch.rand(1, device=device) &lt; 0.5:\n            noise_level = 0.1 + torch.rand(1, device=device).item() * 0.1\n            lam = torch.clamp(tc_part.abs() * noise_level, 0, 5)\n            poisson_noise = torch.poisson(lam)\n            tc_part = tc_part + poisson_noise\n\n    # Spike Band Power (SBP): colored noise\n    if torch.rand(1, device=device) &lt; 0.4:\n        alpha = 0.85 + torch.rand(1, device=device).item() * 0.1\n        noise = torch.randn_like(sbp_part, device=device)\n        filtered_noise = torch.zeros_like(noise)\n\n        for t in range(1, T):\n            filtered_noise[:, t] = alpha * \\\n                filtered_noise[:, t-1] + (1-alpha) * noise[:, t]\n\n        sbp_part = sbp_part + filtered_noise * \\\n            (0.03 + torch.rand(1, device=device).item() * 0.02)\n\n    return torch.cat([tc_part, sbp_part], dim=2)\n\ndef _apply_temporal_masking(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"Apply speech-appropriate temporal masking with smooth edges\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    if torch.rand(1, device=device) &lt; 0.25:\n        max_mask_duration_ms = 200\n        max_mask_bins = max_mask_duration_ms // self.bin_duration_ms\n        num_masks = torch.randint(1, 3, (1,), device=device).item()\n\n        for _ in range(num_masks):\n            mask_duration = torch.randint(\n                1, max_mask_bins + 1, (1,), device=device).item()\n            start_idx = torch.randint(\n                0, max(1, T - mask_duration), (1,), device=device).item()\n\n            mask = torch.ones(T, device=device)\n            mask[start_idx:start_idx + mask_duration] = 0\n\n            window_size = min(2, mask_duration // 2)\n            ramp = torch.linspace(0, 1, steps=window_size+1, device=device)\n\n            for i in range(window_size):\n                if start_idx - i - 1 &gt;= 0:\n                    mask[start_idx - i - 1] = ramp[-(i+2)]\n                if start_idx + mask_duration + i &lt; T:\n                    mask[start_idx + mask_duration + i] = ramp[i+1]\n\n            x = x * mask.view(1, T, 1)\n\n    return x\n\ndef _apply_global_modulation(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"Apply slow global modulation simulating behavioral state changes\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    modulation_freq = 0.05 + torch.rand(1, device=device).item() * 0.1\n    time_vec = torch.arange(T, device=device).float() / 50.0\n\n    tc_modulation = 0.95 + 0.05 * \\\n        torch.sin(2 * np.pi * modulation_freq * time_vec +\n                  torch.rand(1, device=device).item() * 2 * np.pi)\n    sbp_modulation = 0.85 + 0.3 * \\\n        torch.sin(2 * np.pi * modulation_freq * time_vec +\n                  torch.rand(1, device=device).item() * 2 * np.pi)\n\n    x[:, :, self.TC] *= tc_modulation.view(1, T, 1)\n    x[:, :, self.SBP] *= sbp_modulation.view(1, T, 1)\n\n    return x\n\ndef _apply_slow_drift(self, x: torch.Tensor, rank: int = 2, scale: float = 0.01) -&gt; torch.Tensor:\n    \"\"\"Apply low-rank slow drift for cross-session robustness\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    drift_t = torch.randn(B, T, rank, device=device)\n    drift_t = F.avg_pool1d(drift_t.permute(\n        0, 2, 1), kernel_size=25, stride=1, padding=12).permute(0, 2, 1)\n\n    drift_f = torch.randn(rank, C, device=device)\n    drift = torch.matmul(drift_t, drift_f)\n\n    return x + scale * torch.tanh(drift)\n\ndef forward(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"\n    Physiologically-aware augmentation for speech motor cortex BCI data.\n\n    Input:\n        x: torch.Tensor of shape (T, 512) or (B, T, 512)\n\n    Output:\n        Augmented tensor of same shape\n    \"\"\"\n    x, squeeze = self._ensure_batch_dim(x)\n    device = x.device\n\n    # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n    if torch.rand(1, device=device) &lt; 0.35:\n        x = self._temporal_warp(x)\n\n    # 2. Coupled Electrode Dropout (p=0.25) - PHYSIOLOGICALLY ACCURATE\n    if torch.rand(1, device=device) &lt; 0.25:\n        electrode_mask = self._create_coupled_electrode_mask(\n            x, keep_prob=random.uniform(\n                0.75, 0.95))\n        x = x * electrode_mask\n\n    # 3. Array-Level Dropout (p=0.15) - SPATIAL ROBUSTNESS\n    if torch.rand(1, device=device) &lt; 0.15:\n        array_mask = self._create_array_mask(x, dropout_prob=random.uniform(\n            0.1, 3.0))\n        x = x * array_mask\n\n    # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n    if torch.rand(1, device=device) &lt; 0.45:\n        x = self._apply_feature_specific_noise(x)\n\n    # 5. Temporal Masking (p=0.3) - SHORT-TERM ARTIFACTS\n    if torch.rand(1, device=device) &lt; 0.3:\n        x = self._apply_temporal_masking(x)\n\n    # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n    if torch.rand(1, device=device) &lt; 0.2:\n        x = self._apply_global_modulation(x)\n\n    # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n    if torch.rand(1, device=device) &lt; 0.4:\n        x = self._apply_slow_drift(x, rank=random.randint(\n            2, 4), scale=random.uniform(0.005, 0.03))\n    # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n    if torch.rand(1, device=device) &lt; 0.25:\n        x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n            0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n    return self._restore_shape(x, squeeze)\n\ndef forward_tta(self, x: torch.Tensor, num_augments: int = 1, seed: int = None) -&gt; torch.Tensor:\n    \"\"\"\n    Test-Time Augmentation (TTA) method that applies mild, controlled augmentations\n    close to the original data distribution for robust inference.\n\n    Args:\n        x: Input tensor of shape (T, 512) or (B, T, 512)\n        num_augments: Number of augmented versions to generate (default: 1)\n        seed: Optional random seed for reproducibility\n\n    Returns:\n        Augmented tensor. If num_augments &gt; 1, returns tensor of shape (num_augments, B, T, C)\n        otherwise returns same shape as input\n    \"\"\"\n    if seed is not None:\n        torch.manual_seed(seed)\n        if x.device.type == 'cuda':\n            torch.cuda.manual_seed(seed)\n\n    x_original, squeeze = self._ensure_batch_dim(x)\n    B, T, C = x_original.shape\n    device = x_original.device\n\n    # Store multiple augmented versions\n    augmented_versions = []\n\n    for i in range(num_augments):\n        x = x_original.clone()\n\n        # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n        if torch.rand(1, device=device) &lt; 0.35:\n            x = self._temporal_warp(x)\n\n        # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n        if torch.rand(1, device=device) &lt; 0.45:\n            x = self._apply_feature_specific_noise(x)\n\n        # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n        if torch.rand(1, device=device) &lt; 0.2:\n            x = self._apply_global_modulation(x)\n\n        # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n        if torch.rand(1, device=device) &lt; 0.4:\n            x = self._apply_slow_drift(x)\n        # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n        if torch.rand(1, device=device) &lt; 0.25:\n            x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n                0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n        augmented_versions.append(x)\n\n    # Combine results\n    if num_augments == 1:\n        result = augmented_versions[0]\n    else:\n        # Stack along new dimension: (num_augments, B, T, C)\n        result = torch.stack(augmented_versions, dim=0)\n\n    # Restore original shape if needed\n    if num_augments == 1:\n        return self._restore_shape(result, squeeze)\n    else:\n        # For multiple augmentations, keep the extra dimension\n        if squeeze:\n            # If original was (T, C), we now have (num_augments, 1, T, C)\n            # We want to keep the num_augments dimension but remove the batch dim\n            return result.squeeze(1)\n        return result\n</code></pre>\n<p>class SyncTextNeuralAugment:</p>\n<pre><code>\"\"\"\nSynchronized augmentation that randomly cuts segments from both neural signal and text.\nMaintains alignment between neural features and text tokens.\n\"\"\"\n\ndef __init__(self, min_cut_ratio: float = 0.1, max_cut_ratio: float = 0.4, p: float = 0.5):\n    \"\"\"\n    Args:\n        min_cut_ratio: Minimum ratio of sequence to cut out\n        max_cut_ratio: Maximum ratio of sequence to cut out  \n        p: Probability of applying this augmentation\n    \"\"\"\n    self.min_cut_ratio = min_cut_ratio\n    self.max_cut_ratio = max_cut_ratio\n    self.p = p\n\ndef __call__(self, neural: torch.Tensor, text: str) -&gt; Tuple[torch.Tensor, str]:\n    \"\"\"\n    Apply synchronized random cutting to neural signal and text.\n\n    Args:\n        neural: Neural features tensor of shape (T, 512)\n        text: Original text string\n        tokenizer: Tokenizer to help align text with neural segments\n\n    Returns:\n        Cut neural tensor and corresponding text\n    \"\"\"\n    if random.random() &gt; self.p or len(text.strip()) == 0:\n        return neural, text\n\n    T = neural.shape[0]\n    if T &lt; 10:  # Too short to cut meaningfully\n        return neural, text\n\n    # Determine cut parameters\n    cut_ratio = random.uniform(self.min_cut_ratio, self.max_cut_ratio)\n    cut_length = max(1, int(T * cut_ratio))\n    start_idx = random.randint(0, T - cut_length)\n    end_idx = start_idx + cut_length\n\n    # Cut neural signal\n    cut_neural = torch.cat([neural[:start_idx], neural[end_idx:]], dim=0)\n\n    # Cut corresponding text - this is the tricky part\n    # We need to estimate text segments corresponding to neural time steps\n    # Simple approach: assume linear mapping (this may need refinement)\n    text_length = len(text)\n    if text_length &gt; 0:\n        # Calculate text cut positions proportionally\n        text_start_ratio = start_idx / T\n        text_end_ratio = end_idx / T\n\n        text_start_idx = max(0, int(text_length * text_start_ratio))\n        text_end_idx = min(text_length, int(text_length * text_end_ratio))\n\n        # Cut text\n        cut_text = text[:text_start_idx] + text[text_end_idx:]\n    else:\n        cut_text = text\n\n    return cut_neural, cut_text\n</code></pre>\n<p>`</p>",
  "messages": [
    {
      "id": 3384967,
      "postDate": "2026-01-02T10:35:20.533Z",
      "content": "<p>I did not invest much time in this work, particularly in selecting the transformer encoder–decoder, which underperformed compared to RNNs, especially when tokenizing text rather than phonemes. However, I believe I did a solid job on data augmentation, which may be useful for other researchers. here is the code: <a href=\"https://www.kaggle.com/code/jamalsaeedi/b2t-infer\" target=\"_blank\">https://www.kaggle.com/code/jamalsaeedi/b2t-infer</a> </p>\n<p>`</p>\n<p>class PhysioAwareAugment(nn.Module):</p>\n<pre><code>def __init__(self, num_electrodes: int = 256, bin_duration_ms: int = 20):\n    super().__init__()\n    self.num_electrodes = num_electrodes\n    self.bin_duration_ms = bin_duration_ms\n\n    # Array definitions\n    self.arrays = {\n        'ventral_6v': (0, 64),\n        'area_4': (64, 128),\n        'area_55b': (128, 192),\n        'dorsal_6v': (192, 256),\n    }\n\n    self.critical_arrays = ['area_4', 'area_55b']\n\n    # Feature type indices\n    self.TC = slice(0, 256)\n    self.SBP = slice(256, 512)\n\n    # Gaussian smoothing parameters\n    self.smooth_kernel_std = 2.0\n    self.smooth_kernel_size = 100\n    self.has_scipy = HAS_SCIPY\n\ndef _ensure_batch_dim(self, x: torch.Tensor) -&gt; Tuple[torch.Tensor, bool]:\n    \"\"\"Ensure input has batch dimension\"\"\"\n    if x.dim() == 2:  # (T, C)\n        return x.unsqueeze(0), True  # (1, T, C)\n    elif x.dim() == 3:  # (B, T, C)\n        return x, False\n    else:\n        raise ValueError(f\"Expected 2D or 3D input, got shape {x.shape}\")\n\ndef _restore_shape(self, x: torch.Tensor, squeeze: bool) -&gt; torch.Tensor:\n    \"\"\"Restore original shape after augmentation\"\"\"\n    return x.squeeze(0) if squeeze else x\n\ndef _temporal_warp(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"Clean, efficient temporal warping that actually works\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    if T &lt; 10:  # Safety check for short sequences\n        return x\n\n    # Create warp field with 5 control points\n    control_pts = torch.linspace(0, 1, 5, device=device)\n    offsets = (torch.rand(5, device=device) - 0.5) * 0.15\n    warped_pts = torch.clamp(control_pts + offsets, 0, 1)\n    warped_pts, _ = torch.sort(warped_pts)  # Ensure monotonic\n\n    # Interpolate to get full warp field (shape: [T])\n    warp_field = F.interpolate(\n        warped_pts.view(1, 1, -1),\n        size=T,\n        mode='linear',\n        align_corners=True\n    ).squeeze(0).squeeze(0)  # Ensure shape is [T]\n\n    # Map to indices and clamp\n    indices = warp_field * (T - 1)\n    indices = torch.clamp(indices, 0, T - 1)\n\n    # Get floor and ceiling indices\n    floor_idx = torch.floor(indices).long()\n    ceil_idx = torch.ceil(indices).long()\n    ceil_idx = torch.clamp(ceil_idx, 0, T - 1)\n\n    # Interpolation weights (shape: [T])\n    weight = indices - floor_idx.float()\n\n    # Create output tensor\n    x_warped = torch.zeros_like(x)\n\n    # Vectorized interpolation per batch\n    for b in range(B):\n        # For each batch, interpolate all channels at once\n        floor_vals = x[b, floor_idx, :]  # [T, C]\n        ceil_vals = x[b, ceil_idx, :]    # [T, C]\n\n        # Expand weight to match channel dimension: [T] -&gt; [T, 1]\n        weight_expanded = weight.view(-1, 1)\n\n        # Interpolate: (1-weight)*floor + weight*ceil\n        x_warped[b] = floor_vals * \\\n            (1 - weight_expanded) + ceil_vals * weight_expanded\n\n    return x_warped\n\ndef _gauss_smooth(self, inputs: torch.Tensor, smooth_kernel_std: float = None,\n                  smooth_kernel_size: int = None, padding: str = 'same') -&gt; torch.Tensor:\n    \"\"\"\n    Applies 1D Gaussian smoothing with proper fallback if scipy is not available.\n\n    Args:\n        inputs (tensor): B x T x N tensor\n        smooth_kernel_std (float): Standard deviation of Gaussian kernel\n        smooth_kernel_size (int): Size of Gaussian kernel\n        padding (str): Padding mode ('same' or 'valid')\n    \"\"\"\n    if smooth_kernel_std is None:\n        smooth_kernel_std = self.smooth_kernel_std\n    if smooth_kernel_size is None:\n        smooth_kernel_size = self.smooth_kernel_size\n\n    device = inputs.device\n    B, T, C = inputs.shape\n\n    if self.has_scipy:\n        return self._gauss_smooth_scipy(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n    else:\n        return self._gauss_smooth_torch(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n\ndef _gauss_smooth_scipy(self, inputs: torch.Tensor, device: torch.device,\n                        smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                        padding: str = 'same') -&gt; torch.Tensor:\n    \"\"\"Gaussian smoothing using scipy (more accurate)\"\"\"\n    # Get Gaussian kernel\n    inp = np.zeros(smooth_kernel_size, dtype=np.float32)\n    inp[smooth_kernel_size // 2] = 1\n    gaussKernel = gaussian_filter1d(inp, smooth_kernel_std)\n    validIdx = np.argwhere(gaussKernel &gt; 0.01)\n    gaussKernel = gaussKernel[validIdx]\n    gaussKernel = np.squeeze(gaussKernel / np.sum(gaussKernel))\n\n    # Convert to tensor\n    gaussKernel = torch.tensor(\n        gaussKernel, dtype=torch.float32, device=device)\n    gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n    # Prepare convolution\n    B, T, C = inputs.shape\n    inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n    gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n    # Perform convolution\n    smoothed = F.conv1d(inputs_perm, gaussKernel,\n                        padding=padding, groups=C)\n    return smoothed.permute(0, 2, 1)  # [B, T, C]\n\ndef _gauss_smooth_torch(self, inputs: torch.Tensor, device: torch.device,\n                        smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                        padding: str = 'same') -&gt; torch.Tensor:\n    \"\"\"Gaussian smoothing using pure PyTorch (fallback)\"\"\"\n    # Create Gaussian kernel using PyTorch\n    x = torch.arange(-(smooth_kernel_size//2), smooth_kernel_size //\n                     2 + 1, device=device, dtype=torch.float32)\n    gaussKernel = torch.exp(-x**2 / (2 * smooth_kernel_std**2))\n    gaussKernel = gaussKernel / gaussKernel.sum()\n\n    # Ensure kernel has odd size\n    if gaussKernel.size(0) % 2 == 0:\n        gaussKernel = gaussKernel[:-1]\n\n    gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n    # Prepare convolution\n    B, T, C = inputs.shape\n    inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n    gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n    # Handle padding\n    if padding == 'same':\n        padding_size = gaussKernel.size(2) // 2\n    else:\n        padding_size = 0\n\n    # Perform convolution\n    smoothed = F.conv1d(inputs_perm, gaussKernel,\n                        padding=padding_size, groups=C)\n    return smoothed.permute(0, 2, 1)  # [B, T, C]\n\ndef _create_coupled_electrode_mask(self, x: torch.Tensor, keep_prob: float = 0.85) -&gt; torch.Tensor:\n    \"\"\"Create mask where both TC and SBP features for same electrode are dropped together\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    # Base mask for 256 electrodes\n    elec_mask = torch.bernoulli(torch.ones(\n        self.num_electrodes, device=device) * keep_prob)\n\n    # Add spatial correlation within arrays\n    for array_name, (start_idx, end_idx) in self.arrays.items():\n        center_elec = torch.randint(\n            start_idx, end_idx, (1,), device=device).item()\n        distances = torch.arange(\n            start_idx, end_idx, device=device).float() - center_elec\n        spatial_weights = torch.exp(-torch.abs(distances) / 10.0)\n\n        spatial_bias = 0.25 * spatial_weights\n        p = torch.clamp(keep_prob - spatial_bias, 0.05, 1.0)\n\n        array_mask = torch.bernoulli(p)\n        elec_mask[start_idx:end_idx] = array_mask\n\n    # Expand to 512 features: [TC_mask, SBP_mask]\n    full_mask = torch.cat([elec_mask, elec_mask])\n    return full_mask.view(1, 1, -1)  # (1, 1, 512)\n\ndef _create_array_mask(self, x: torch.Tensor, dropout_prob: float = 0.1) -&gt; torch.Tensor:\n    \"\"\"Create mask that drops entire arrays with anatomical awareness\"\"\"\n    B, T, C = x.shape\n    device = x.device\n    mask = torch.ones(512, device=device)\n\n    for array_name, (start_idx, end_idx) in self.arrays.items():\n        is_critical = array_name in self.critical_arrays\n        array_dropout_prob = dropout_prob * 0.3 if is_critical else dropout_prob\n\n        if torch.rand(1, device=device) &lt; array_dropout_prob:\n            mask[start_idx:end_idx] = 0  # TC features\n            mask[start_idx + 256:end_idx + 256] = 0  # SBP features\n\n    return mask.view(1, 1, -1)  # (1, 1, 512)\n\ndef _apply_feature_specific_noise(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"Apply physiologically appropriate noise to different feature types\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    tc_part = x[:, :, self.TC].clone()\n    sbp_part = x[:, :, self.SBP].clone()\n\n    # Threshold Crossings (TC): multiplicative noise\n    if torch.rand(1, device=device) &lt; 0.3:\n        scale_factor = 0.9 + torch.rand(1, device=device).item() * 0.2\n        tc_part = tc_part * scale_factor\n\n        if torch.rand(1, device=device) &lt; 0.5:\n            noise_level = 0.1 + torch.rand(1, device=device).item() * 0.1\n            lam = torch.clamp(tc_part.abs() * noise_level, 0, 5)\n            poisson_noise = torch.poisson(lam)\n            tc_part = tc_part + poisson_noise\n\n    # Spike Band Power (SBP): colored noise\n    if torch.rand(1, device=device) &lt; 0.4:\n        alpha = 0.85 + torch.rand(1, device=device).item() * 0.1\n        noise = torch.randn_like(sbp_part, device=device)\n        filtered_noise = torch.zeros_like(noise)\n\n        for t in range(1, T):\n            filtered_noise[:, t] = alpha * \\\n                filtered_noise[:, t-1] + (1-alpha) * noise[:, t]\n\n        sbp_part = sbp_part + filtered_noise * \\\n            (0.03 + torch.rand(1, device=device).item() * 0.02)\n\n    return torch.cat([tc_part, sbp_part], dim=2)\n\ndef _apply_temporal_masking(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"Apply speech-appropriate temporal masking with smooth edges\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    if torch.rand(1, device=device) &lt; 0.25:\n        max_mask_duration_ms = 200\n        max_mask_bins = max_mask_duration_ms // self.bin_duration_ms\n        num_masks = torch.randint(1, 3, (1,), device=device).item()\n\n        for _ in range(num_masks):\n            mask_duration = torch.randint(\n                1, max_mask_bins + 1, (1,), device=device).item()\n            start_idx = torch.randint(\n                0, max(1, T - mask_duration), (1,), device=device).item()\n\n            mask = torch.ones(T, device=device)\n            mask[start_idx:start_idx + mask_duration] = 0\n\n            window_size = min(2, mask_duration // 2)\n            ramp = torch.linspace(0, 1, steps=window_size+1, device=device)\n\n            for i in range(window_size):\n                if start_idx - i - 1 &gt;= 0:\n                    mask[start_idx - i - 1] = ramp[-(i+2)]\n                if start_idx + mask_duration + i &lt; T:\n                    mask[start_idx + mask_duration + i] = ramp[i+1]\n\n            x = x * mask.view(1, T, 1)\n\n    return x\n\ndef _apply_global_modulation(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"Apply slow global modulation simulating behavioral state changes\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    modulation_freq = 0.05 + torch.rand(1, device=device).item() * 0.1\n    time_vec = torch.arange(T, device=device).float() / 50.0\n\n    tc_modulation = 0.95 + 0.05 * \\\n        torch.sin(2 * np.pi * modulation_freq * time_vec +\n                  torch.rand(1, device=device).item() * 2 * np.pi)\n    sbp_modulation = 0.85 + 0.3 * \\\n        torch.sin(2 * np.pi * modulation_freq * time_vec +\n                  torch.rand(1, device=device).item() * 2 * np.pi)\n\n    x[:, :, self.TC] *= tc_modulation.view(1, T, 1)\n    x[:, :, self.SBP] *= sbp_modulation.view(1, T, 1)\n\n    return x\n\ndef _apply_slow_drift(self, x: torch.Tensor, rank: int = 2, scale: float = 0.01) -&gt; torch.Tensor:\n    \"\"\"Apply low-rank slow drift for cross-session robustness\"\"\"\n    B, T, C = x.shape\n    device = x.device\n\n    drift_t = torch.randn(B, T, rank, device=device)\n    drift_t = F.avg_pool1d(drift_t.permute(\n        0, 2, 1), kernel_size=25, stride=1, padding=12).permute(0, 2, 1)\n\n    drift_f = torch.randn(rank, C, device=device)\n    drift = torch.matmul(drift_t, drift_f)\n\n    return x + scale * torch.tanh(drift)\n\ndef forward(self, x: torch.Tensor) -&gt; torch.Tensor:\n    \"\"\"\n    Physiologically-aware augmentation for speech motor cortex BCI data.\n\n    Input:\n        x: torch.Tensor of shape (T, 512) or (B, T, 512)\n\n    Output:\n        Augmented tensor of same shape\n    \"\"\"\n    x, squeeze = self._ensure_batch_dim(x)\n    device = x.device\n\n    # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n    if torch.rand(1, device=device) &lt; 0.35:\n        x = self._temporal_warp(x)\n\n    # 2. Coupled Electrode Dropout (p=0.25) - PHYSIOLOGICALLY ACCURATE\n    if torch.rand(1, device=device) &lt; 0.25:\n        electrode_mask = self._create_coupled_electrode_mask(\n            x, keep_prob=random.uniform(\n                0.75, 0.95))\n        x = x * electrode_mask\n\n    # 3. Array-Level Dropout (p=0.15) - SPATIAL ROBUSTNESS\n    if torch.rand(1, device=device) &lt; 0.15:\n        array_mask = self._create_array_mask(x, dropout_prob=random.uniform(\n            0.1, 3.0))\n        x = x * array_mask\n\n    # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n    if torch.rand(1, device=device) &lt; 0.45:\n        x = self._apply_feature_specific_noise(x)\n\n    # 5. Temporal Masking (p=0.3) - SHORT-TERM ARTIFACTS\n    if torch.rand(1, device=device) &lt; 0.3:\n        x = self._apply_temporal_masking(x)\n\n    # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n    if torch.rand(1, device=device) &lt; 0.2:\n        x = self._apply_global_modulation(x)\n\n    # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n    if torch.rand(1, device=device) &lt; 0.4:\n        x = self._apply_slow_drift(x, rank=random.randint(\n            2, 4), scale=random.uniform(0.005, 0.03))\n    # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n    if torch.rand(1, device=device) &lt; 0.25:\n        x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n            0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n    return self._restore_shape(x, squeeze)\n\ndef forward_tta(self, x: torch.Tensor, num_augments: int = 1, seed: int = None) -&gt; torch.Tensor:\n    \"\"\"\n    Test-Time Augmentation (TTA) method that applies mild, controlled augmentations\n    close to the original data distribution for robust inference.\n\n    Args:\n        x: Input tensor of shape (T, 512) or (B, T, 512)\n        num_augments: Number of augmented versions to generate (default: 1)\n        seed: Optional random seed for reproducibility\n\n    Returns:\n        Augmented tensor. If num_augments &gt; 1, returns tensor of shape (num_augments, B, T, C)\n        otherwise returns same shape as input\n    \"\"\"\n    if seed is not None:\n        torch.manual_seed(seed)\n        if x.device.type == 'cuda':\n            torch.cuda.manual_seed(seed)\n\n    x_original, squeeze = self._ensure_batch_dim(x)\n    B, T, C = x_original.shape\n    device = x_original.device\n\n    # Store multiple augmented versions\n    augmented_versions = []\n\n    for i in range(num_augments):\n        x = x_original.clone()\n\n        # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n        if torch.rand(1, device=device) &lt; 0.35:\n            x = self._temporal_warp(x)\n\n        # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n        if torch.rand(1, device=device) &lt; 0.45:\n            x = self._apply_feature_specific_noise(x)\n\n        # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n        if torch.rand(1, device=device) &lt; 0.2:\n            x = self._apply_global_modulation(x)\n\n        # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n        if torch.rand(1, device=device) &lt; 0.4:\n            x = self._apply_slow_drift(x)\n        # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n        if torch.rand(1, device=device) &lt; 0.25:\n            x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n                0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n        augmented_versions.append(x)\n\n    # Combine results\n    if num_augments == 1:\n        result = augmented_versions[0]\n    else:\n        # Stack along new dimension: (num_augments, B, T, C)\n        result = torch.stack(augmented_versions, dim=0)\n\n    # Restore original shape if needed\n    if num_augments == 1:\n        return self._restore_shape(result, squeeze)\n    else:\n        # For multiple augmentations, keep the extra dimension\n        if squeeze:\n            # If original was (T, C), we now have (num_augments, 1, T, C)\n            # We want to keep the num_augments dimension but remove the batch dim\n            return result.squeeze(1)\n        return result\n</code></pre>\n<p>class SyncTextNeuralAugment:</p>\n<pre><code>\"\"\"\nSynchronized augmentation that randomly cuts segments from both neural signal and text.\nMaintains alignment between neural features and text tokens.\n\"\"\"\n\ndef __init__(self, min_cut_ratio: float = 0.1, max_cut_ratio: float = 0.4, p: float = 0.5):\n    \"\"\"\n    Args:\n        min_cut_ratio: Minimum ratio of sequence to cut out\n        max_cut_ratio: Maximum ratio of sequence to cut out  \n        p: Probability of applying this augmentation\n    \"\"\"\n    self.min_cut_ratio = min_cut_ratio\n    self.max_cut_ratio = max_cut_ratio\n    self.p = p\n\ndef __call__(self, neural: torch.Tensor, text: str) -&gt; Tuple[torch.Tensor, str]:\n    \"\"\"\n    Apply synchronized random cutting to neural signal and text.\n\n    Args:\n        neural: Neural features tensor of shape (T, 512)\n        text: Original text string\n        tokenizer: Tokenizer to help align text with neural segments\n\n    Returns:\n        Cut neural tensor and corresponding text\n    \"\"\"\n    if random.random() &gt; self.p or len(text.strip()) == 0:\n        return neural, text\n\n    T = neural.shape[0]\n    if T &lt; 10:  # Too short to cut meaningfully\n        return neural, text\n\n    # Determine cut parameters\n    cut_ratio = random.uniform(self.min_cut_ratio, self.max_cut_ratio)\n    cut_length = max(1, int(T * cut_ratio))\n    start_idx = random.randint(0, T - cut_length)\n    end_idx = start_idx + cut_length\n\n    # Cut neural signal\n    cut_neural = torch.cat([neural[:start_idx], neural[end_idx:]], dim=0)\n\n    # Cut corresponding text - this is the tricky part\n    # We need to estimate text segments corresponding to neural time steps\n    # Simple approach: assume linear mapping (this may need refinement)\n    text_length = len(text)\n    if text_length &gt; 0:\n        # Calculate text cut positions proportionally\n        text_start_ratio = start_idx / T\n        text_end_ratio = end_idx / T\n\n        text_start_idx = max(0, int(text_length * text_start_ratio))\n        text_end_idx = min(text_length, int(text_length * text_end_ratio))\n\n        # Cut text\n        cut_text = text[:text_start_idx] + text[text_end_idx:]\n    else:\n        cut_text = text\n\n    return cut_neural, cut_text\n</code></pre>\n<p>`</p>",
      "rawMarkdown": "I did not invest much time in this work, particularly in selecting the transformer encoder–decoder, which underperformed compared to RNNs, especially when tokenizing text rather than phonemes. However, I believe I did a solid job on data augmentation, which may be useful for other researchers. here is the code: https://www.kaggle.com/code/jamalsaeedi/b2t-infer \n\n`\n\nclass PhysioAwareAugment(nn.Module):\n\n    def __init__(self, num_electrodes: int = 256, bin_duration_ms: int = 20):\n        super().__init__()\n        self.num_electrodes = num_electrodes\n        self.bin_duration_ms = bin_duration_ms\n\n        # Array definitions\n        self.arrays = {\n            'ventral_6v': (0, 64),\n            'area_4': (64, 128),\n            'area_55b': (128, 192),\n            'dorsal_6v': (192, 256),\n        }\n\n        self.critical_arrays = ['area_4', 'area_55b']\n\n        # Feature type indices\n        self.TC = slice(0, 256)\n        self.SBP = slice(256, 512)\n\n        # Gaussian smoothing parameters\n        self.smooth_kernel_std = 2.0\n        self.smooth_kernel_size = 100\n        self.has_scipy = HAS_SCIPY\n\n    def _ensure_batch_dim(self, x: torch.Tensor) -> Tuple[torch.Tensor, bool]:\n        \"\"\"Ensure input has batch dimension\"\"\"\n        if x.dim() == 2:  # (T, C)\n            return x.unsqueeze(0), True  # (1, T, C)\n        elif x.dim() == 3:  # (B, T, C)\n            return x, False\n        else:\n            raise ValueError(f\"Expected 2D or 3D input, got shape {x.shape}\")\n\n    def _restore_shape(self, x: torch.Tensor, squeeze: bool) -> torch.Tensor:\n        \"\"\"Restore original shape after augmentation\"\"\"\n        return x.squeeze(0) if squeeze else x\n\n    def _temporal_warp(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Clean, efficient temporal warping that actually works\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        if T < 10:  # Safety check for short sequences\n            return x\n\n        # Create warp field with 5 control points\n        control_pts = torch.linspace(0, 1, 5, device=device)\n        offsets = (torch.rand(5, device=device) - 0.5) * 0.15\n        warped_pts = torch.clamp(control_pts + offsets, 0, 1)\n        warped_pts, _ = torch.sort(warped_pts)  # Ensure monotonic\n\n        # Interpolate to get full warp field (shape: [T])\n        warp_field = F.interpolate(\n            warped_pts.view(1, 1, -1),\n            size=T,\n            mode='linear',\n            align_corners=True\n        ).squeeze(0).squeeze(0)  # Ensure shape is [T]\n\n        # Map to indices and clamp\n        indices = warp_field * (T - 1)\n        indices = torch.clamp(indices, 0, T - 1)\n\n        # Get floor and ceiling indices\n        floor_idx = torch.floor(indices).long()\n        ceil_idx = torch.ceil(indices).long()\n        ceil_idx = torch.clamp(ceil_idx, 0, T - 1)\n\n        # Interpolation weights (shape: [T])\n        weight = indices - floor_idx.float()\n\n        # Create output tensor\n        x_warped = torch.zeros_like(x)\n\n        # Vectorized interpolation per batch\n        for b in range(B):\n            # For each batch, interpolate all channels at once\n            floor_vals = x[b, floor_idx, :]  # [T, C]\n            ceil_vals = x[b, ceil_idx, :]    # [T, C]\n\n            # Expand weight to match channel dimension: [T] -> [T, 1]\n            weight_expanded = weight.view(-1, 1)\n\n            # Interpolate: (1-weight)*floor + weight*ceil\n            x_warped[b] = floor_vals * \\\n                (1 - weight_expanded) + ceil_vals * weight_expanded\n\n        return x_warped\n\n    def _gauss_smooth(self, inputs: torch.Tensor, smooth_kernel_std: float = None,\n                      smooth_kernel_size: int = None, padding: str = 'same') -> torch.Tensor:\n        \"\"\"\n        Applies 1D Gaussian smoothing with proper fallback if scipy is not available.\n\n        Args:\n            inputs (tensor): B x T x N tensor\n            smooth_kernel_std (float): Standard deviation of Gaussian kernel\n            smooth_kernel_size (int): Size of Gaussian kernel\n            padding (str): Padding mode ('same' or 'valid')\n        \"\"\"\n        if smooth_kernel_std is None:\n            smooth_kernel_std = self.smooth_kernel_std\n        if smooth_kernel_size is None:\n            smooth_kernel_size = self.smooth_kernel_size\n\n        device = inputs.device\n        B, T, C = inputs.shape\n\n        if self.has_scipy:\n            return self._gauss_smooth_scipy(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n        else:\n            return self._gauss_smooth_torch(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n\n    def _gauss_smooth_scipy(self, inputs: torch.Tensor, device: torch.device,\n                            smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                            padding: str = 'same') -> torch.Tensor:\n        \"\"\"Gaussian smoothing using scipy (more accurate)\"\"\"\n        # Get Gaussian kernel\n        inp = np.zeros(smooth_kernel_size, dtype=np.float32)\n        inp[smooth_kernel_size // 2] = 1\n        gaussKernel = gaussian_filter1d(inp, smooth_kernel_std)\n        validIdx = np.argwhere(gaussKernel > 0.01)\n        gaussKernel = gaussKernel[validIdx]\n        gaussKernel = np.squeeze(gaussKernel / np.sum(gaussKernel))\n\n        # Convert to tensor\n        gaussKernel = torch.tensor(\n            gaussKernel, dtype=torch.float32, device=device)\n        gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n        # Prepare convolution\n        B, T, C = inputs.shape\n        inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n        gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n        # Perform convolution\n        smoothed = F.conv1d(inputs_perm, gaussKernel,\n                            padding=padding, groups=C)\n        return smoothed.permute(0, 2, 1)  # [B, T, C]\n\n    def _gauss_smooth_torch(self, inputs: torch.Tensor, device: torch.device,\n                            smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                            padding: str = 'same') -> torch.Tensor:\n        \"\"\"Gaussian smoothing using pure PyTorch (fallback)\"\"\"\n        # Create Gaussian kernel using PyTorch\n        x = torch.arange(-(smooth_kernel_size//2), smooth_kernel_size //\n                         2 + 1, device=device, dtype=torch.float32)\n        gaussKernel = torch.exp(-x**2 / (2 * smooth_kernel_std**2))\n        gaussKernel = gaussKernel / gaussKernel.sum()\n\n        # Ensure kernel has odd size\n        if gaussKernel.size(0) % 2 == 0:\n            gaussKernel = gaussKernel[:-1]\n\n        gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n        # Prepare convolution\n        B, T, C = inputs.shape\n        inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n        gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n        # Handle padding\n        if padding == 'same':\n            padding_size = gaussKernel.size(2) // 2\n        else:\n            padding_size = 0\n\n        # Perform convolution\n        smoothed = F.conv1d(inputs_perm, gaussKernel,\n                            padding=padding_size, groups=C)\n        return smoothed.permute(0, 2, 1)  # [B, T, C]\n\n    def _create_coupled_electrode_mask(self, x: torch.Tensor, keep_prob: float = 0.85) -> torch.Tensor:\n        \"\"\"Create mask where both TC and SBP features for same electrode are dropped together\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        # Base mask for 256 electrodes\n        elec_mask = torch.bernoulli(torch.ones(\n            self.num_electrodes, device=device) * keep_prob)\n\n        # Add spatial correlation within arrays\n        for array_name, (start_idx, end_idx) in self.arrays.items():\n            center_elec = torch.randint(\n                start_idx, end_idx, (1,), device=device).item()\n            distances = torch.arange(\n                start_idx, end_idx, device=device).float() - center_elec\n            spatial_weights = torch.exp(-torch.abs(distances) / 10.0)\n\n            spatial_bias = 0.25 * spatial_weights\n            p = torch.clamp(keep_prob - spatial_bias, 0.05, 1.0)\n\n            array_mask = torch.bernoulli(p)\n            elec_mask[start_idx:end_idx] = array_mask\n\n        # Expand to 512 features: [TC_mask, SBP_mask]\n        full_mask = torch.cat([elec_mask, elec_mask])\n        return full_mask.view(1, 1, -1)  # (1, 1, 512)\n\n    def _create_array_mask(self, x: torch.Tensor, dropout_prob: float = 0.1) -> torch.Tensor:\n        \"\"\"Create mask that drops entire arrays with anatomical awareness\"\"\"\n        B, T, C = x.shape\n        device = x.device\n        mask = torch.ones(512, device=device)\n\n        for array_name, (start_idx, end_idx) in self.arrays.items():\n            is_critical = array_name in self.critical_arrays\n            array_dropout_prob = dropout_prob * 0.3 if is_critical else dropout_prob\n\n            if torch.rand(1, device=device) < array_dropout_prob:\n                mask[start_idx:end_idx] = 0  # TC features\n                mask[start_idx + 256:end_idx + 256] = 0  # SBP features\n\n        return mask.view(1, 1, -1)  # (1, 1, 512)\n\n    def _apply_feature_specific_noise(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply physiologically appropriate noise to different feature types\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        tc_part = x[:, :, self.TC].clone()\n        sbp_part = x[:, :, self.SBP].clone()\n\n        # Threshold Crossings (TC): multiplicative noise\n        if torch.rand(1, device=device) < 0.3:\n            scale_factor = 0.9 + torch.rand(1, device=device).item() * 0.2\n            tc_part = tc_part * scale_factor\n\n            if torch.rand(1, device=device) < 0.5:\n                noise_level = 0.1 + torch.rand(1, device=device).item() * 0.1\n                lam = torch.clamp(tc_part.abs() * noise_level, 0, 5)\n                poisson_noise = torch.poisson(lam)\n                tc_part = tc_part + poisson_noise\n\n        # Spike Band Power (SBP): colored noise\n        if torch.rand(1, device=device) < 0.4:\n            alpha = 0.85 + torch.rand(1, device=device).item() * 0.1\n            noise = torch.randn_like(sbp_part, device=device)\n            filtered_noise = torch.zeros_like(noise)\n\n            for t in range(1, T):\n                filtered_noise[:, t] = alpha * \\\n                    filtered_noise[:, t-1] + (1-alpha) * noise[:, t]\n\n            sbp_part = sbp_part + filtered_noise * \\\n                (0.03 + torch.rand(1, device=device).item() * 0.02)\n\n        return torch.cat([tc_part, sbp_part], dim=2)\n\n    def _apply_temporal_masking(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply speech-appropriate temporal masking with smooth edges\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        if torch.rand(1, device=device) < 0.25:\n            max_mask_duration_ms = 200\n            max_mask_bins = max_mask_duration_ms // self.bin_duration_ms\n            num_masks = torch.randint(1, 3, (1,), device=device).item()\n\n            for _ in range(num_masks):\n                mask_duration = torch.randint(\n                    1, max_mask_bins + 1, (1,), device=device).item()\n                start_idx = torch.randint(\n                    0, max(1, T - mask_duration), (1,), device=device).item()\n\n                mask = torch.ones(T, device=device)\n                mask[start_idx:start_idx + mask_duration] = 0\n\n                window_size = min(2, mask_duration // 2)\n                ramp = torch.linspace(0, 1, steps=window_size+1, device=device)\n\n                for i in range(window_size):\n                    if start_idx - i - 1 >= 0:\n                        mask[start_idx - i - 1] = ramp[-(i+2)]\n                    if start_idx + mask_duration + i < T:\n                        mask[start_idx + mask_duration + i] = ramp[i+1]\n\n                x = x * mask.view(1, T, 1)\n\n        return x\n\n    def _apply_global_modulation(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply slow global modulation simulating behavioral state changes\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        modulation_freq = 0.05 + torch.rand(1, device=device).item() * 0.1\n        time_vec = torch.arange(T, device=device).float() / 50.0\n\n        tc_modulation = 0.95 + 0.05 * \\\n            torch.sin(2 * np.pi * modulation_freq * time_vec +\n                      torch.rand(1, device=device).item() * 2 * np.pi)\n        sbp_modulation = 0.85 + 0.3 * \\\n            torch.sin(2 * np.pi * modulation_freq * time_vec +\n                      torch.rand(1, device=device).item() * 2 * np.pi)\n\n        x[:, :, self.TC] *= tc_modulation.view(1, T, 1)\n        x[:, :, self.SBP] *= sbp_modulation.view(1, T, 1)\n\n        return x\n\n    def _apply_slow_drift(self, x: torch.Tensor, rank: int = 2, scale: float = 0.01) -> torch.Tensor:\n        \"\"\"Apply low-rank slow drift for cross-session robustness\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        drift_t = torch.randn(B, T, rank, device=device)\n        drift_t = F.avg_pool1d(drift_t.permute(\n            0, 2, 1), kernel_size=25, stride=1, padding=12).permute(0, 2, 1)\n\n        drift_f = torch.randn(rank, C, device=device)\n        drift = torch.matmul(drift_t, drift_f)\n\n        return x + scale * torch.tanh(drift)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Physiologically-aware augmentation for speech motor cortex BCI data.\n\n        Input:\n            x: torch.Tensor of shape (T, 512) or (B, T, 512)\n\n        Output:\n            Augmented tensor of same shape\n        \"\"\"\n        x, squeeze = self._ensure_batch_dim(x)\n        device = x.device\n\n        # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n        if torch.rand(1, device=device) < 0.35:\n            x = self._temporal_warp(x)\n\n        # 2. Coupled Electrode Dropout (p=0.25) - PHYSIOLOGICALLY ACCURATE\n        if torch.rand(1, device=device) < 0.25:\n            electrode_mask = self._create_coupled_electrode_mask(\n                x, keep_prob=random.uniform(\n                    0.75, 0.95))\n            x = x * electrode_mask\n\n        # 3. Array-Level Dropout (p=0.15) - SPATIAL ROBUSTNESS\n        if torch.rand(1, device=device) < 0.15:\n            array_mask = self._create_array_mask(x, dropout_prob=random.uniform(\n                0.1, 3.0))\n            x = x * array_mask\n\n        # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n        if torch.rand(1, device=device) < 0.45:\n            x = self._apply_feature_specific_noise(x)\n\n        # 5. Temporal Masking (p=0.3) - SHORT-TERM ARTIFACTS\n        if torch.rand(1, device=device) < 0.3:\n            x = self._apply_temporal_masking(x)\n\n        # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n        if torch.rand(1, device=device) < 0.2:\n            x = self._apply_global_modulation(x)\n\n        # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n        if torch.rand(1, device=device) < 0.4:\n            x = self._apply_slow_drift(x, rank=random.randint(\n                2, 4), scale=random.uniform(0.005, 0.03))\n        # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n        if torch.rand(1, device=device) < 0.25:\n            x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n                0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n        return self._restore_shape(x, squeeze)\n\n    def forward_tta(self, x: torch.Tensor, num_augments: int = 1, seed: int = None) -> torch.Tensor:\n        \"\"\"\n        Test-Time Augmentation (TTA) method that applies mild, controlled augmentations\n        close to the original data distribution for robust inference.\n\n        Args:\n            x: Input tensor of shape (T, 512) or (B, T, 512)\n            num_augments: Number of augmented versions to generate (default: 1)\n            seed: Optional random seed for reproducibility\n\n        Returns:\n            Augmented tensor. If num_augments > 1, returns tensor of shape (num_augments, B, T, C)\n            otherwise returns same shape as input\n        \"\"\"\n        if seed is not None:\n            torch.manual_seed(seed)\n            if x.device.type == 'cuda':\n                torch.cuda.manual_seed(seed)\n\n        x_original, squeeze = self._ensure_batch_dim(x)\n        B, T, C = x_original.shape\n        device = x_original.device\n\n        # Store multiple augmented versions\n        augmented_versions = []\n\n        for i in range(num_augments):\n            x = x_original.clone()\n\n            # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n            if torch.rand(1, device=device) < 0.35:\n                x = self._temporal_warp(x)\n\n            # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n            if torch.rand(1, device=device) < 0.45:\n                x = self._apply_feature_specific_noise(x)\n\n            # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n            if torch.rand(1, device=device) < 0.2:\n                x = self._apply_global_modulation(x)\n\n            # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n            if torch.rand(1, device=device) < 0.4:\n                x = self._apply_slow_drift(x)\n            # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n            if torch.rand(1, device=device) < 0.25:\n                x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n                    0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n            augmented_versions.append(x)\n\n        # Combine results\n        if num_augments == 1:\n            result = augmented_versions[0]\n        else:\n            # Stack along new dimension: (num_augments, B, T, C)\n            result = torch.stack(augmented_versions, dim=0)\n\n        # Restore original shape if needed\n        if num_augments == 1:\n            return self._restore_shape(result, squeeze)\n        else:\n            # For multiple augmentations, keep the extra dimension\n            if squeeze:\n                # If original was (T, C), we now have (num_augments, 1, T, C)\n                # We want to keep the num_augments dimension but remove the batch dim\n                return result.squeeze(1)\n            return result\n\n\nclass SyncTextNeuralAugment:\n\n    \"\"\"\n    Synchronized augmentation that randomly cuts segments from both neural signal and text.\n    Maintains alignment between neural features and text tokens.\n    \"\"\"\n\n    def __init__(self, min_cut_ratio: float = 0.1, max_cut_ratio: float = 0.4, p: float = 0.5):\n        \"\"\"\n        Args:\n            min_cut_ratio: Minimum ratio of sequence to cut out\n            max_cut_ratio: Maximum ratio of sequence to cut out  \n            p: Probability of applying this augmentation\n        \"\"\"\n        self.min_cut_ratio = min_cut_ratio\n        self.max_cut_ratio = max_cut_ratio\n        self.p = p\n\n    def __call__(self, neural: torch.Tensor, text: str) -> Tuple[torch.Tensor, str]:\n        \"\"\"\n        Apply synchronized random cutting to neural signal and text.\n\n        Args:\n            neural: Neural features tensor of shape (T, 512)\n            text: Original text string\n            tokenizer: Tokenizer to help align text with neural segments\n\n        Returns:\n            Cut neural tensor and corresponding text\n        \"\"\"\n        if random.random() > self.p or len(text.strip()) == 0:\n            return neural, text\n\n        T = neural.shape[0]\n        if T < 10:  # Too short to cut meaningfully\n            return neural, text\n\n        # Determine cut parameters\n        cut_ratio = random.uniform(self.min_cut_ratio, self.max_cut_ratio)\n        cut_length = max(1, int(T * cut_ratio))\n        start_idx = random.randint(0, T - cut_length)\n        end_idx = start_idx + cut_length\n\n        # Cut neural signal\n        cut_neural = torch.cat([neural[:start_idx], neural[end_idx:]], dim=0)\n\n        # Cut corresponding text - this is the tricky part\n        # We need to estimate text segments corresponding to neural time steps\n        # Simple approach: assume linear mapping (this may need refinement)\n        text_length = len(text)\n        if text_length > 0:\n            # Calculate text cut positions proportionally\n            text_start_ratio = start_idx / T\n            text_end_ratio = end_idx / T\n\n            text_start_idx = max(0, int(text_length * text_start_ratio))\n            text_end_idx = min(text_length, int(text_length * text_end_ratio))\n\n            # Cut text\n            cut_text = text[:text_start_idx] + text[text_end_idx:]\n        else:\n            cut_text = text\n\n        return cut_neural, cut_text\n`"
    }
  ],
  "comments": [],
  "raw_markdown_by_id": {
    "3384967": "I did not invest much time in this work, particularly in selecting the transformer encoder–decoder, which underperformed compared to RNNs, especially when tokenizing text rather than phonemes. However, I believe I did a solid job on data augmentation, which may be useful for other researchers. here is the code: https://www.kaggle.com/code/jamalsaeedi/b2t-infer \n\n`\n\nclass PhysioAwareAugment(nn.Module):\n\n    def __init__(self, num_electrodes: int = 256, bin_duration_ms: int = 20):\n        super().__init__()\n        self.num_electrodes = num_electrodes\n        self.bin_duration_ms = bin_duration_ms\n\n        # Array definitions\n        self.arrays = {\n            'ventral_6v': (0, 64),\n            'area_4': (64, 128),\n            'area_55b': (128, 192),\n            'dorsal_6v': (192, 256),\n        }\n\n        self.critical_arrays = ['area_4', 'area_55b']\n\n        # Feature type indices\n        self.TC = slice(0, 256)\n        self.SBP = slice(256, 512)\n\n        # Gaussian smoothing parameters\n        self.smooth_kernel_std = 2.0\n        self.smooth_kernel_size = 100\n        self.has_scipy = HAS_SCIPY\n\n    def _ensure_batch_dim(self, x: torch.Tensor) -> Tuple[torch.Tensor, bool]:\n        \"\"\"Ensure input has batch dimension\"\"\"\n        if x.dim() == 2:  # (T, C)\n            return x.unsqueeze(0), True  # (1, T, C)\n        elif x.dim() == 3:  # (B, T, C)\n            return x, False\n        else:\n            raise ValueError(f\"Expected 2D or 3D input, got shape {x.shape}\")\n\n    def _restore_shape(self, x: torch.Tensor, squeeze: bool) -> torch.Tensor:\n        \"\"\"Restore original shape after augmentation\"\"\"\n        return x.squeeze(0) if squeeze else x\n\n    def _temporal_warp(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Clean, efficient temporal warping that actually works\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        if T < 10:  # Safety check for short sequences\n            return x\n\n        # Create warp field with 5 control points\n        control_pts = torch.linspace(0, 1, 5, device=device)\n        offsets = (torch.rand(5, device=device) - 0.5) * 0.15\n        warped_pts = torch.clamp(control_pts + offsets, 0, 1)\n        warped_pts, _ = torch.sort(warped_pts)  # Ensure monotonic\n\n        # Interpolate to get full warp field (shape: [T])\n        warp_field = F.interpolate(\n            warped_pts.view(1, 1, -1),\n            size=T,\n            mode='linear',\n            align_corners=True\n        ).squeeze(0).squeeze(0)  # Ensure shape is [T]\n\n        # Map to indices and clamp\n        indices = warp_field * (T - 1)\n        indices = torch.clamp(indices, 0, T - 1)\n\n        # Get floor and ceiling indices\n        floor_idx = torch.floor(indices).long()\n        ceil_idx = torch.ceil(indices).long()\n        ceil_idx = torch.clamp(ceil_idx, 0, T - 1)\n\n        # Interpolation weights (shape: [T])\n        weight = indices - floor_idx.float()\n\n        # Create output tensor\n        x_warped = torch.zeros_like(x)\n\n        # Vectorized interpolation per batch\n        for b in range(B):\n            # For each batch, interpolate all channels at once\n            floor_vals = x[b, floor_idx, :]  # [T, C]\n            ceil_vals = x[b, ceil_idx, :]    # [T, C]\n\n            # Expand weight to match channel dimension: [T] -> [T, 1]\n            weight_expanded = weight.view(-1, 1)\n\n            # Interpolate: (1-weight)*floor + weight*ceil\n            x_warped[b] = floor_vals * \\\n                (1 - weight_expanded) + ceil_vals * weight_expanded\n\n        return x_warped\n\n    def _gauss_smooth(self, inputs: torch.Tensor, smooth_kernel_std: float = None,\n                      smooth_kernel_size: int = None, padding: str = 'same') -> torch.Tensor:\n        \"\"\"\n        Applies 1D Gaussian smoothing with proper fallback if scipy is not available.\n\n        Args:\n            inputs (tensor): B x T x N tensor\n            smooth_kernel_std (float): Standard deviation of Gaussian kernel\n            smooth_kernel_size (int): Size of Gaussian kernel\n            padding (str): Padding mode ('same' or 'valid')\n        \"\"\"\n        if smooth_kernel_std is None:\n            smooth_kernel_std = self.smooth_kernel_std\n        if smooth_kernel_size is None:\n            smooth_kernel_size = self.smooth_kernel_size\n\n        device = inputs.device\n        B, T, C = inputs.shape\n\n        if self.has_scipy:\n            return self._gauss_smooth_scipy(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n        else:\n            return self._gauss_smooth_torch(inputs, device, smooth_kernel_std, smooth_kernel_size, padding)\n\n    def _gauss_smooth_scipy(self, inputs: torch.Tensor, device: torch.device,\n                            smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                            padding: str = 'same') -> torch.Tensor:\n        \"\"\"Gaussian smoothing using scipy (more accurate)\"\"\"\n        # Get Gaussian kernel\n        inp = np.zeros(smooth_kernel_size, dtype=np.float32)\n        inp[smooth_kernel_size // 2] = 1\n        gaussKernel = gaussian_filter1d(inp, smooth_kernel_std)\n        validIdx = np.argwhere(gaussKernel > 0.01)\n        gaussKernel = gaussKernel[validIdx]\n        gaussKernel = np.squeeze(gaussKernel / np.sum(gaussKernel))\n\n        # Convert to tensor\n        gaussKernel = torch.tensor(\n            gaussKernel, dtype=torch.float32, device=device)\n        gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n        # Prepare convolution\n        B, T, C = inputs.shape\n        inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n        gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n        # Perform convolution\n        smoothed = F.conv1d(inputs_perm, gaussKernel,\n                            padding=padding, groups=C)\n        return smoothed.permute(0, 2, 1)  # [B, T, C]\n\n    def _gauss_smooth_torch(self, inputs: torch.Tensor, device: torch.device,\n                            smooth_kernel_std: float = 2.0, smooth_kernel_size: int = 100,\n                            padding: str = 'same') -> torch.Tensor:\n        \"\"\"Gaussian smoothing using pure PyTorch (fallback)\"\"\"\n        # Create Gaussian kernel using PyTorch\n        x = torch.arange(-(smooth_kernel_size//2), smooth_kernel_size //\n                         2 + 1, device=device, dtype=torch.float32)\n        gaussKernel = torch.exp(-x**2 / (2 * smooth_kernel_std**2))\n        gaussKernel = gaussKernel / gaussKernel.sum()\n\n        # Ensure kernel has odd size\n        if gaussKernel.size(0) % 2 == 0:\n            gaussKernel = gaussKernel[:-1]\n\n        gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n        # Prepare convolution\n        B, T, C = inputs.shape\n        inputs_perm = inputs.permute(0, 2, 1)  # [B, C, T]\n        gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n        # Handle padding\n        if padding == 'same':\n            padding_size = gaussKernel.size(2) // 2\n        else:\n            padding_size = 0\n\n        # Perform convolution\n        smoothed = F.conv1d(inputs_perm, gaussKernel,\n                            padding=padding_size, groups=C)\n        return smoothed.permute(0, 2, 1)  # [B, T, C]\n\n    def _create_coupled_electrode_mask(self, x: torch.Tensor, keep_prob: float = 0.85) -> torch.Tensor:\n        \"\"\"Create mask where both TC and SBP features for same electrode are dropped together\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        # Base mask for 256 electrodes\n        elec_mask = torch.bernoulli(torch.ones(\n            self.num_electrodes, device=device) * keep_prob)\n\n        # Add spatial correlation within arrays\n        for array_name, (start_idx, end_idx) in self.arrays.items():\n            center_elec = torch.randint(\n                start_idx, end_idx, (1,), device=device).item()\n            distances = torch.arange(\n                start_idx, end_idx, device=device).float() - center_elec\n            spatial_weights = torch.exp(-torch.abs(distances) / 10.0)\n\n            spatial_bias = 0.25 * spatial_weights\n            p = torch.clamp(keep_prob - spatial_bias, 0.05, 1.0)\n\n            array_mask = torch.bernoulli(p)\n            elec_mask[start_idx:end_idx] = array_mask\n\n        # Expand to 512 features: [TC_mask, SBP_mask]\n        full_mask = torch.cat([elec_mask, elec_mask])\n        return full_mask.view(1, 1, -1)  # (1, 1, 512)\n\n    def _create_array_mask(self, x: torch.Tensor, dropout_prob: float = 0.1) -> torch.Tensor:\n        \"\"\"Create mask that drops entire arrays with anatomical awareness\"\"\"\n        B, T, C = x.shape\n        device = x.device\n        mask = torch.ones(512, device=device)\n\n        for array_name, (start_idx, end_idx) in self.arrays.items():\n            is_critical = array_name in self.critical_arrays\n            array_dropout_prob = dropout_prob * 0.3 if is_critical else dropout_prob\n\n            if torch.rand(1, device=device) < array_dropout_prob:\n                mask[start_idx:end_idx] = 0  # TC features\n                mask[start_idx + 256:end_idx + 256] = 0  # SBP features\n\n        return mask.view(1, 1, -1)  # (1, 1, 512)\n\n    def _apply_feature_specific_noise(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply physiologically appropriate noise to different feature types\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        tc_part = x[:, :, self.TC].clone()\n        sbp_part = x[:, :, self.SBP].clone()\n\n        # Threshold Crossings (TC): multiplicative noise\n        if torch.rand(1, device=device) < 0.3:\n            scale_factor = 0.9 + torch.rand(1, device=device).item() * 0.2\n            tc_part = tc_part * scale_factor\n\n            if torch.rand(1, device=device) < 0.5:\n                noise_level = 0.1 + torch.rand(1, device=device).item() * 0.1\n                lam = torch.clamp(tc_part.abs() * noise_level, 0, 5)\n                poisson_noise = torch.poisson(lam)\n                tc_part = tc_part + poisson_noise\n\n        # Spike Band Power (SBP): colored noise\n        if torch.rand(1, device=device) < 0.4:\n            alpha = 0.85 + torch.rand(1, device=device).item() * 0.1\n            noise = torch.randn_like(sbp_part, device=device)\n            filtered_noise = torch.zeros_like(noise)\n\n            for t in range(1, T):\n                filtered_noise[:, t] = alpha * \\\n                    filtered_noise[:, t-1] + (1-alpha) * noise[:, t]\n\n            sbp_part = sbp_part + filtered_noise * \\\n                (0.03 + torch.rand(1, device=device).item() * 0.02)\n\n        return torch.cat([tc_part, sbp_part], dim=2)\n\n    def _apply_temporal_masking(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply speech-appropriate temporal masking with smooth edges\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        if torch.rand(1, device=device) < 0.25:\n            max_mask_duration_ms = 200\n            max_mask_bins = max_mask_duration_ms // self.bin_duration_ms\n            num_masks = torch.randint(1, 3, (1,), device=device).item()\n\n            for _ in range(num_masks):\n                mask_duration = torch.randint(\n                    1, max_mask_bins + 1, (1,), device=device).item()\n                start_idx = torch.randint(\n                    0, max(1, T - mask_duration), (1,), device=device).item()\n\n                mask = torch.ones(T, device=device)\n                mask[start_idx:start_idx + mask_duration] = 0\n\n                window_size = min(2, mask_duration // 2)\n                ramp = torch.linspace(0, 1, steps=window_size+1, device=device)\n\n                for i in range(window_size):\n                    if start_idx - i - 1 >= 0:\n                        mask[start_idx - i - 1] = ramp[-(i+2)]\n                    if start_idx + mask_duration + i < T:\n                        mask[start_idx + mask_duration + i] = ramp[i+1]\n\n                x = x * mask.view(1, T, 1)\n\n        return x\n\n    def _apply_global_modulation(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply slow global modulation simulating behavioral state changes\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        modulation_freq = 0.05 + torch.rand(1, device=device).item() * 0.1\n        time_vec = torch.arange(T, device=device).float() / 50.0\n\n        tc_modulation = 0.95 + 0.05 * \\\n            torch.sin(2 * np.pi * modulation_freq * time_vec +\n                      torch.rand(1, device=device).item() * 2 * np.pi)\n        sbp_modulation = 0.85 + 0.3 * \\\n            torch.sin(2 * np.pi * modulation_freq * time_vec +\n                      torch.rand(1, device=device).item() * 2 * np.pi)\n\n        x[:, :, self.TC] *= tc_modulation.view(1, T, 1)\n        x[:, :, self.SBP] *= sbp_modulation.view(1, T, 1)\n\n        return x\n\n    def _apply_slow_drift(self, x: torch.Tensor, rank: int = 2, scale: float = 0.01) -> torch.Tensor:\n        \"\"\"Apply low-rank slow drift for cross-session robustness\"\"\"\n        B, T, C = x.shape\n        device = x.device\n\n        drift_t = torch.randn(B, T, rank, device=device)\n        drift_t = F.avg_pool1d(drift_t.permute(\n            0, 2, 1), kernel_size=25, stride=1, padding=12).permute(0, 2, 1)\n\n        drift_f = torch.randn(rank, C, device=device)\n        drift = torch.matmul(drift_t, drift_f)\n\n        return x + scale * torch.tanh(drift)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Physiologically-aware augmentation for speech motor cortex BCI data.\n\n        Input:\n            x: torch.Tensor of shape (T, 512) or (B, T, 512)\n\n        Output:\n            Augmented tensor of same shape\n        \"\"\"\n        x, squeeze = self._ensure_batch_dim(x)\n        device = x.device\n\n        # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n        if torch.rand(1, device=device) < 0.35:\n            x = self._temporal_warp(x)\n\n        # 2. Coupled Electrode Dropout (p=0.25) - PHYSIOLOGICALLY ACCURATE\n        if torch.rand(1, device=device) < 0.25:\n            electrode_mask = self._create_coupled_electrode_mask(\n                x, keep_prob=random.uniform(\n                    0.75, 0.95))\n            x = x * electrode_mask\n\n        # 3. Array-Level Dropout (p=0.15) - SPATIAL ROBUSTNESS\n        if torch.rand(1, device=device) < 0.15:\n            array_mask = self._create_array_mask(x, dropout_prob=random.uniform(\n                0.1, 3.0))\n            x = x * array_mask\n\n        # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n        if torch.rand(1, device=device) < 0.45:\n            x = self._apply_feature_specific_noise(x)\n\n        # 5. Temporal Masking (p=0.3) - SHORT-TERM ARTIFACTS\n        if torch.rand(1, device=device) < 0.3:\n            x = self._apply_temporal_masking(x)\n\n        # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n        if torch.rand(1, device=device) < 0.2:\n            x = self._apply_global_modulation(x)\n\n        # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n        if torch.rand(1, device=device) < 0.4:\n            x = self._apply_slow_drift(x, rank=random.randint(\n                2, 4), scale=random.uniform(0.005, 0.03))\n        # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n        if torch.rand(1, device=device) < 0.25:\n            x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n                0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n        return self._restore_shape(x, squeeze)\n\n    def forward_tta(self, x: torch.Tensor, num_augments: int = 1, seed: int = None) -> torch.Tensor:\n        \"\"\"\n        Test-Time Augmentation (TTA) method that applies mild, controlled augmentations\n        close to the original data distribution for robust inference.\n\n        Args:\n            x: Input tensor of shape (T, 512) or (B, T, 512)\n            num_augments: Number of augmented versions to generate (default: 1)\n            seed: Optional random seed for reproducibility\n\n        Returns:\n            Augmented tensor. If num_augments > 1, returns tensor of shape (num_augments, B, T, C)\n            otherwise returns same shape as input\n        \"\"\"\n        if seed is not None:\n            torch.manual_seed(seed)\n            if x.device.type == 'cuda':\n                torch.cuda.manual_seed(seed)\n\n        x_original, squeeze = self._ensure_batch_dim(x)\n        B, T, C = x_original.shape\n        device = x_original.device\n\n        # Store multiple augmented versions\n        augmented_versions = []\n\n        for i in range(num_augments):\n            x = x_original.clone()\n\n            # 1. Temporal Warping (p=0.35) - SPEECH TEMPORAL VARIABILITY\n            if torch.rand(1, device=device) < 0.35:\n                x = self._temporal_warp(x)\n\n            # 4. Feature-Specific Noise (p=0.45) - PHYSIOLOGICAL NOISE MODELS\n            if torch.rand(1, device=device) < 0.45:\n                x = self._apply_feature_specific_noise(x)\n\n            # 6. Global Signal Modulation (p=0.2) - BEHAVIORAL STATE VARIATIONS\n            if torch.rand(1, device=device) < 0.2:\n                x = self._apply_global_modulation(x)\n\n            # 7. Slow Drift (p=0.4) - CROSS-SESSION ROBUSTNESS\n            if torch.rand(1, device=device) < 0.4:\n                x = self._apply_slow_drift(x)\n            # 8. Gaussian Smoothing (p=0.25) - SIMULATE LOW-PASS FILTERING EFFECTS\n            if torch.rand(1, device=device) < 0.25:\n                x = self._gauss_smooth(x, smooth_kernel_std=random.uniform(\n                    0.3, 2.0), smooth_kernel_size=random.choice([9, 15, 21, 27, 35]))\n\n            augmented_versions.append(x)\n\n        # Combine results\n        if num_augments == 1:\n            result = augmented_versions[0]\n        else:\n            # Stack along new dimension: (num_augments, B, T, C)\n            result = torch.stack(augmented_versions, dim=0)\n\n        # Restore original shape if needed\n        if num_augments == 1:\n            return self._restore_shape(result, squeeze)\n        else:\n            # For multiple augmentations, keep the extra dimension\n            if squeeze:\n                # If original was (T, C), we now have (num_augments, 1, T, C)\n                # We want to keep the num_augments dimension but remove the batch dim\n                return result.squeeze(1)\n            return result\n\n\nclass SyncTextNeuralAugment:\n\n    \"\"\"\n    Synchronized augmentation that randomly cuts segments from both neural signal and text.\n    Maintains alignment between neural features and text tokens.\n    \"\"\"\n\n    def __init__(self, min_cut_ratio: float = 0.1, max_cut_ratio: float = 0.4, p: float = 0.5):\n        \"\"\"\n        Args:\n            min_cut_ratio: Minimum ratio of sequence to cut out\n            max_cut_ratio: Maximum ratio of sequence to cut out  \n            p: Probability of applying this augmentation\n        \"\"\"\n        self.min_cut_ratio = min_cut_ratio\n        self.max_cut_ratio = max_cut_ratio\n        self.p = p\n\n    def __call__(self, neural: torch.Tensor, text: str) -> Tuple[torch.Tensor, str]:\n        \"\"\"\n        Apply synchronized random cutting to neural signal and text.\n\n        Args:\n            neural: Neural features tensor of shape (T, 512)\n            text: Original text string\n            tokenizer: Tokenizer to help align text with neural segments\n\n        Returns:\n            Cut neural tensor and corresponding text\n        \"\"\"\n        if random.random() > self.p or len(text.strip()) == 0:\n            return neural, text\n\n        T = neural.shape[0]\n        if T < 10:  # Too short to cut meaningfully\n            return neural, text\n\n        # Determine cut parameters\n        cut_ratio = random.uniform(self.min_cut_ratio, self.max_cut_ratio)\n        cut_length = max(1, int(T * cut_ratio))\n        start_idx = random.randint(0, T - cut_length)\n        end_idx = start_idx + cut_length\n\n        # Cut neural signal\n        cut_neural = torch.cat([neural[:start_idx], neural[end_idx:]], dim=0)\n\n        # Cut corresponding text - this is the tricky part\n        # We need to estimate text segments corresponding to neural time steps\n        # Simple approach: assume linear mapping (this may need refinement)\n        text_length = len(text)\n        if text_length > 0:\n            # Calculate text cut positions proportionally\n            text_start_ratio = start_idx / T\n            text_end_ratio = end_idx / T\n\n            text_start_idx = max(0, int(text_length * text_start_ratio))\n            text_end_idx = min(text_length, int(text_length * text_end_ratio))\n\n            # Cut text\n            cut_text = text[:text_start_idx] + text[text_end_idx:]\n        else:\n            cut_text = text\n\n        return cut_neural, cut_text\n`"
  }
}