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)# pip install -U transformers accelerate # 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
File size: 9,723 Bytes
c61906f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 | """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")
|