| name | meta-learning-ict-brain-decoding |
| description | 元学习上下文方法实现无需训练的跨被试脑解码。通过上下文学习实现训练无关的跨个体fMRI解码。适用于零样本脑解码、快速脑机接口、个体化神经科学。触发词:元学习脑解码、上下文学习、跨被试、训练无关、零样本。 |
Meta-Learning In-Context for Brain Decoding
元学习上下文(Meta-Learning In-Context)方法实现无需训练(Training-Free)的跨被试脑解码。
Metadata
- Source: arXiv:2604.08537
- Title: Meta-learning In-Context Enables Training-Free Cross Subject Brain Decoding
- Authors: Mu Nan, Muquan Yu, Weijian Mai, Jacob S. Prince, Hossein Adeli, Rui Zhang, Jiahang Cao, Benjamin Becker, John A. Pyles, Margaret M. Henderson, Chunfeng Song, Nikolaus Kriegeskorte, Michael J. Tarr, Xiaoqing Hu, Andrew F. Luo
- Published: 2026-04-09
- Category: Neuroscience & Machine Learning
Core Methodology
Key Innovation
该研究提出了Meta-Learning In-Context (ML-IC)框架,利用元学习的上下文学习能力,在测试时通过少量上下文示例直接适应新被试,无需额外的训练或微调。这解决了传统脑解码中跨被试泛化性差和需要大量训练数据的问题。
Technical Framework
-
元学习上下文(ML-IC)机制
- 在训练阶段学习跨被试的通用解码策略
- 测试时提供目标被试的少量示例作为"上下文"
- 模型自动适应新被试的特征空间
-
训练无关适应
- 无需反向传播或参数更新
- 完全基于前向传播的上下文推理
- 支持在线快速适应
-
跨被试泛化
- 学习被试无关的神经表征
- 处理个体间神经变异
- 零样本或少样本迁移
Implementation Guide
Prerequisites
- Python 3.8+
- PyTorch 1.9+
- nilearn, nibabel
- 预训练模型支持
Step-by-Step
-
数据准备
import nibabel as nib
from nilearn import datasets, input_data
hcp_dataset = datasets.fetch_hcp(...)
masker = input_data.NiftiLabelsMasker(
labels_img='Schaefer2018_400Parcels_7Networks_order_FSLMNI152_1mm.nii.gz',
standardize=True
)
-
上下文构建
def build_context_set(target_subject, support_subjects, n_context=20):
"""
构建上下文示例集
Args:
target_subject: 目标被试ID
support_subjects: 支持被试列表
n_context: 上下文示例数量
Returns:
context_trials: 上下文试次
context_labels: 上下文标签
"""
context_trials = []
context_labels = []
for subject in support_subjects[:n_context]:
trial = load_subject_trial(subject)
label = get_trial_label(subject, trial)
context_trials.append(trial)
context_labels.append(label)
return np.array(context_trials), np.array(context_labels)
-
ML-IC解码器
import torch
import torch.nn as nn
class MLICBrainDecoder(nn.Module):
def __init__(
self,
input_dim=400,
hidden_dim=512,
output_dim=10,
n_heads=8
):
super().__init__()
self.encoder = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.LayerNorm(hidden_dim),
nn.ReLU(),
nn.Dropout(0.1),
nn.Linear(hidden_dim, hidden_dim)
)
self.cross_attention = nn.MultiheadAttention(
hidden_dim, n_heads, batch_first=True
)
self.decoder = nn.Sequential(
nn.Linear(hidden_dim * 2, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, output_dim)
)
def forward(self, query_fmri, context_fmri, context_labels):
"""
Args:
query_fmri: 目标试次 [batch, input_dim]
context_fmri: 上下文fMRI [n_context, input_dim]
context_labels: 上下文标签 [n_context]
Returns:
logits: 预测 logits [batch, output_dim]
"""
query_encoded = self.encoder(query_fmri)
context_encoded = self.encoder(context_fmri)
query_expanded = query_encoded.unsqueeze()
context_expanded = context_encoded.unsqueeze().expand(
query_fmri.size(), -, -
)
attended, attention_weights = .cross_attention(
query_expanded,
context_expanded,
context_expanded
)
query_attended = attended.squeeze()
fused = torch.cat([query_encoded, query_attended], dim=-)
logits = .decoder(fused)
logits, attention_weights
-
元学习训练
def meta_train_step(model, batch, optimizer):
"""
元学习训练步骤 (MAML风格)
"""
support_fmri, support_labels, query_fmri, query_labels = batch
logits, _ = model(query_fmri, support_fmri, support_labels)
loss = nn.CrossEntropyLoss()(logits, query_labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss.item()
Code Example
"""
元学习上下文脑解码实现
"""
import torch
import torch.nn as nn
import numpy as np
from typing import List, Tuple, Dict
from torch.utils.data import Dataset, DataLoader
import nibabel as nib
class BrainDecodingDataset(Dataset):
"""脑解码数据集"""
def __init__(self, fmri_data, labels, subject_ids):
self.fmri_data = torch.FloatTensor(fmri_data)
self.labels = torch.LongTensor(labels)
self.subject_ids = subject_ids
self.subjects = list(set(subject_ids))
def __len__(self):
return len(self.fmri_data)
def __getitem__(self, idx):
return {
'fmri': self.fmri_data[idx],
'label': self.labels[idx],
'subject': self.subject_ids[idx]
}
def get_episode(self, n_support, n_query, subjects=):
subjects :
subjects = .subjects
selected_subject = np.random.choice(subjects)
subject_mask = np.array(.subject_ids) == selected_subject
subject_indices = np.where(subject_mask)[]
sampled = np.random.choice(
subject_indices,
size=n_support + n_query,
replace=
)
support_idx = sampled[:n_support]
query_idx = sampled[n_support:]
(
.fmri_data[support_idx],
.labels[support_idx],
.fmri_data[query_idx],
.labels[query_idx]
)
(nn.Module):
():
().__init__()
.encoder = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.LayerNorm(hidden_dim),
nn.ReLU(),
nn.Dropout(),
nn.Linear(hidden_dim, hidden_dim)
)
.attention = nn.MultiheadAttention(
hidden_dim, n_heads, batch_first=
)
.adaptive_norm = nn.LayerNorm(hidden_dim)
.classifier = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Dropout(),
nn.Linear(hidden_dim, output_dim)
)
():
.encoder(x)
() -> [torch.Tensor, torch.Tensor]:
query_enc = .encode(query)
context_enc = .encode(context)
batch_size = query.size()
n_context = context.size()
query_expanded = query_enc.unsqueeze()
context_expanded = context_enc.unsqueeze().expand(
batch_size, n_context, -
)
attended, attn_weights = .attention(
query_expanded,
context_expanded,
context_expanded
)
output = query_expanded + attended
output = .adaptive_norm(output)
output = output.squeeze()
logits = .classifier(output)
logits, attn_weights
() -> torch.Tensor:
torch.no_grad():
logits, _ = .forward(
new_subject_fmri,
context_fmri,
context_labels
)
predictions = torch.argmax(logits, dim=-)
predictions
:
():
.model = model.to(device)
.optimizer = optimizer
.device = device
.criterion = nn.CrossEntropyLoss()
():
.model.train()
episode (n_episodes):
support_fmri, support_labels, query_fmri, query_labels = \
dataset.get_episode(n_support, n_query)
support_fmri = support_fmri.to(.device)
support_labels = support_labels.to(.device)
query_fmri = query_fmri.to(.device)
query_labels = query_labels.to(.device)
logits, _ = .model(
query_fmri,
support_fmri,
support_labels
)
loss = .criterion(logits, query_labels)
.optimizer.zero_grad()
loss.backward()
.optimizer.step()
episode % == :
acc = (logits.argmax(dim=-) == query_labels).().mean()
()
() -> [, ]:
.model.()
context_subjects = [s s train_dataset.subjects s != test_subject]
context_fmri_list = []
context_labels_list = []
_ (n_context):
subj = np.random.choice(context_subjects)
subj_mask = np.array(train_dataset.subject_ids) == subj
subj_indices = np.where(subj_mask)[]
idx = np.random.choice(subj_indices)
context_fmri_list.append(train_dataset.fmri_data[idx])
context_labels_list.append(train_dataset.labels[idx])
context_fmri = torch.stack(context_fmri_list).to(.device)
context_labels = torch.tensor(context_labels_list).to(.device)
test_fmri = test_fmri.to(.device)
test_labels = test_labels.to(.device)
torch.no_grad():
logits, _ = .model(
test_fmri,
context_fmri,
context_labels
)
predictions = logits.argmax(dim=-)
accuracy = (predictions == test_labels).().mean().item()
{
: accuracy,
: (test_labels),
: n_context
}
():
n_subjects =
n_trials_per_subject =
n_regions =
n_classes =
fmri_data = []
labels = []
subject_ids = []
subject (n_subjects):
trial (n_trials_per_subject):
subject_pattern = np.random.randn(n_regions) *
trial_pattern = np.random.randn(n_regions) *
fmri = subject_pattern + trial_pattern + np.random.randn(n_regions) *
fmri_data.append(fmri)
labels.append(trial % n_classes)
subject_ids.append()
fmri_data = np.array(fmri_data)
labels = np.array(labels)
dataset = BrainDecodingDataset(fmri_data, labels, subject_ids)
model = MetaLearner(
input_dim=n_regions,
hidden_dim=,
output_dim=n_classes
)
optimizer = torch.optim.Adam(model.parameters(), lr=)
trainer = BrainDecodingTrainer(model, optimizer, device=)
()
trainer.train_episode(dataset, n_episodes=)
()
new_subject_fmri = torch.randn(, n_regions)
new_subject_labels = torch.randint(, n_classes, (,))
results = trainer.evaluate_cross_subject(
dataset,
,
new_subject_fmri,
new_subject_labels,
n_context=
)
()
__name__ == :
main()
Applications
-
快速脑机接口(BCI)
- 无需训练的新型用户适配
- 实时脑信号解码
- 消费级BCI设备
-
临床神经科学
-
认知神经科学研究
-
脑解码基准测试
Pitfalls
- 上下文选择: 上下文示例质量影响性能
- 被试变异: 极端个体差异可能降低效果
- 任务限制: 可能局限于特定类型任务
- 计算成本: 推理时需要存储上下文
- 数据质量: 依赖预训练数据质量
Related Skills
- computational-lesions-multilingual-language-models
- vlm-visual-cortex-alignment-robustness
- sensorless-gaze-following-neuroscience
- brain-dit-fmri-foundation-model
References
- arXiv:2604.08537 (2026)
- MAML: Model-Agnostic Meta-Learning (Finn et al.)
- Learning to Learn (Thrun & Pratt)
- Neural Processes (Garnelo et al.)