MANAS-2 / model.py
mannasAI-labs's picture
Release MANAS-2 encoder
c61906f verified
Raw History Blame Contribute Delete
9.72 kB
"""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")