Feature Extraction
Transformers
Safetensors
PyTorch
English
manas2
eeg
neuroscience
foundation-model
masked-autoencoder
representation-learning
custom_code
Instructions to use MannasAI/MANAS-2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use MannasAI/MANAS-2 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="MannasAI/MANAS-2", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("MannasAI/MANAS-2", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download model.py from MannasAI/MANAS-2: direct link, hf CLI and curl.
- Browser
- Download file 9.72 kB
-
https://huggingface.co/MannasAI/MANAS-2/resolve/main/model.py
- Command line
-
hf download hf://MannasAI/MANAS-2/model.py
-
curl -L -o model.py https://huggingface.co/MannasAI/MANAS-2/resolve/main/model.py
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 | |
| 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.""" | |
| 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") | |