"""MANAS-2 EEG encoder (200 Hz, 1-second patches). model = Manas2Model.from_pretrained(".") features = model(eeg, positions) # [batch, patches, channels, 512] EEG: preprocessed float32 [B, C, T]. Positions: float32 [B, C, 3] in centimeters, matching the training coordinate frame and channel order. Training decoders and losses are not needed for feature extraction. """ from __future__ import annotations import json import math from pathlib import Path import torch from torch import nn from torch.nn import functional as F from transformers import PretrainedConfig, PreTrainedModel class Manas2Config(PretrainedConfig): """Configuration for the released MANAS-2 encoder.""" model_type = "manas2" def __init__(self, hidden_size=512, num_hidden_layers=22, num_attention_heads=8, patch_size=200, patch_stride=180, sampling_rate=200, electrode_positions=None, **kwargs): super().__init__(**kwargs) values = (hidden_size, num_hidden_layers, num_attention_heads, patch_size, patch_stride, sampling_rate) if values != (512, 22, 8, 200, 180, 200): raise ValueError("This release implements only the 22-layer MANAS-2 architecture") self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.patch_size = patch_size self.patch_stride = patch_stride self.sampling_rate = sampling_rate # Original training coordinate table in meters; lookup converts to cm. self.electrode_positions = electrode_positions or {} class RMSNorm(nn.Module): def __init__(self, dim): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + 1e-8) * self.weight class GEGLU(nn.Module): def forward(self, x): value, gate = x.chunk(2, dim=-1) # Preserve the checkpoint's original gate, including both terms. gelu = 0.5 * (1.0 + torch.tanh(math.sqrt(2.0 / math.pi) * (gate + 0.044715 * gate.pow(3)))) return value * (gate * torch.sigmoid(gate) + gate * gelu) class FeedForward(nn.Module): def __init__(self, dim): super().__init__() self.in_proj = nn.Linear(dim, 8 * dim) self.gate = GEGLU() self.out_proj = nn.Linear(4 * dim, dim) def forward(self, x): return self.out_proj(self.gate(self.in_proj(x))) class Attention(nn.Module): def __init__(self, dim, heads): super().__init__() self.heads = heads self.qkv_proj = nn.Linear(dim, 3 * dim) self.out_proj = nn.Linear(dim, dim) def forward(self, x): batch, tokens, dim = x.shape q, k, v = self.qkv_proj(x).chunk(3, dim=-1) q, k, v = [t.view(batch, tokens, self.heads, dim // self.heads).transpose(1, 2) for t in (q, k, v)] x = F.scaled_dot_product_attention(q, k, v, dropout_p=0.0) return self.out_proj(x.transpose(1, 2).contiguous().view(batch, tokens, dim)) class Block(nn.Module): def __init__(self, dim, heads): super().__init__() self.pre_attn_norm = RMSNorm(dim) self.attn = Attention(dim, heads) self.pre_ffn_norm = RMSNorm(dim) self.ffn = FeedForward(dim) def forward(self, x): x = x + self.attn(self.pre_attn_norm(x)) return x + self.ffn(self.pre_ffn_norm(x)) class Encoder(nn.Module): def __init__(self, dim=512, depth=22, heads=8): super().__init__() self.layers = nn.ModuleList([Block(dim, heads) for _ in range(depth)]) self.final_norm = RMSNorm(dim) def forward(self, x): for layer in self.layers: x = layer(x) return self.final_norm(x) class PatchEmbed(nn.Module): def __init__(self): super().__init__() self.linear = nn.Linear(200, 512, bias=False) def forward(self, eeg): return self.linear(eeg.unfold(-1, 200, 180)) class PosEnc(nn.Module): def __init__(self): super().__init__() freqs = torch.linspace(1.0, 10.0, 4) self.register_buffer("freq_matrix", torch.cartesian_prod(*([freqs] * 4)).T) self.fourier_linear = nn.Linear(512, 512, bias=False) self.learned_linear = nn.Sequential(nn.Linear(4, 1024, bias=False), GEGLU(), RMSNorm(512)) self.final_norm = RMSNorm(512) def forward(self, coords): # HF loads contiguous buffers; preserve the original Fourier matrix layout. phases = coords @ self.freq_matrix.T.contiguous().T fourier = self.fourier_linear(torch.cat([phases.sin(), phases.cos()], dim=-1)) return self.final_norm(fourier + self.learned_linear(coords)) def channel_positions(names): """Look up exact channel names in the training table; return [C, 3] in cm.""" if not names: raise ValueError("At least one channel is required") table = json.loads(Path(__file__).with_name("positions.json").read_text()) return torch.tensor([table[name] for name in names], dtype=torch.float32) * 100.0 class Manas2Model(PreTrainedModel): """22-layer pretrained encoder; forward returns [B, patches, C, 512]. Patches are 200 samples with a 180-sample stride. Incomplete trailing patches are dropped. No normalization, resampling, or pooling is applied. """ config_class = Manas2Config base_model_prefix = "" main_input_name = "eeg" _no_split_modules = ["Block"] patch_size = 200 step = 180 embed_dim = 512 def __init__(self, config=None): super().__init__(config if config is not None else Manas2Config()) self.patch_embed = PatchEmbed() self.pos_enc = PosEnc() self.encoder = Encoder() self.post_init() def get_channel_positions(self, names): """Return exact-name training coordinates [C, 3] in centimeters on CPU.""" if not names: raise ValueError("At least one channel is required") table = self.config.electrode_positions if not table: raise ValueError("No electrode table in config; pass coordinates explicitly") unknown = [name for name in names if name not in table] if unknown: raise ValueError(f"Unknown electrode names: {unknown}; supply matching coordinates") return torch.tensor([table[name] for name in names], dtype=torch.float32) * 100.0 @classmethod def from_checkpoint(cls, path=".", *, device="cpu"): """Load a local release directory, Safetensors file, or original .pt. Encoder tensors load strictly. Only known training-only tensors are discarded from the full checkpoint. Encoder-only state dicts also work. """ path = Path(path).expanduser() if path.is_dir(): path = path / "model.safetensors" if path.suffix == ".safetensors": from safetensors.torch import load_file state = load_file(str(path)) else: state = torch.load(path, map_location="cpu", weights_only=True) state = state.get("model_state_dict", state.get("state_dict", state)) for prefix in ("module.", "wrapped_model.", "model.", "hybrid."): if state and all(key.startswith(prefix) for key in state): state = {key[len(prefix):]: value for key, value in state.items()} encoder_state = {} training_prefixes = ("decoder.", "band_decoder.", "band_target_module.", "band_pooler.", "aux_linear.", "aux_predict.") for key, value in state.items(): if key.startswith(("patch_embed.", "pos_enc.", "encoder.")): encoder_state[key] = value elif key != "aux_query" and not key.startswith(training_prefixes): raise ValueError(f"Unexpected checkpoint tensor: {key}") table_path = Path(__file__).with_name("positions.json") table = json.loads(table_path.read_text()) if table_path.exists() else {} model = cls(Manas2Config(electrode_positions=table)) model.load_state_dict(encoder_state, strict=True) return model.to(device).eval() def forward(self, eeg, positions): if eeg.ndim != 3 or min(eeg.shape[:2]) < 1 or eeg.shape[-1] < self.patch_size: raise ValueError("EEG must have shape [B, C, T] with B,C >= 1 and T >= 200") if tuple(positions.shape) != (*eeg.shape[:2], 3): raise ValueError("Positions must have shape [B, C, 3]") batch, channels, _ = eeg.shape tokens = self.patch_embed(eeg) patches = tokens.shape[2] spatial = positions.unsqueeze(2).expand(-1, -1, patches, -1) time = torch.arange(patches, device=eeg.device, dtype=torch.float32) time = time.view(1, 1, patches, 1).expand(batch, channels, -1, -1) coords = torch.cat([spatial, time], dim=-1).reshape(batch, channels * patches, 4) coords = coords.to(dtype=self.pos_enc.freq_matrix.dtype) encoded = self.encoder(tokens.flatten(1, 2) + self.pos_enc(coords)) return encoded.reshape(batch, channels, patches, 512).permute(0, 2, 1, 3).contiguous() def export_features(self, eeg, positions): return self(eeg, positions) class ConRecRBH(Manas2Model): """Compatibility name for the original local-only encoder interface.""" @classmethod def from_pretrained(cls, path=".", *, device="cpu"): return cls.from_checkpoint(path, device=device) Manas2Config.register_for_auto_class() Manas2Model.register_for_auto_class("AutoModel")