import os import re import torch import gradio as gr import numpy as np import nibabel as nib from pathlib import Path from dataclasses import dataclass from typing import Dict, List, Tuple, Optional import torch.nn as nn import torch.nn.functional as F from einops import rearrange from einops.layers.torch import Rearrange from scipy.ndimage import zoom import matplotlib.pyplot as plt import seaborn as sns # core config @dataclass class Config: VOLUME_SIZE: Tuple[int, int, int] = (64, 64, 30) EMBED_DIM: int = 256 NUM_HEADS: int = 8 NUM_LAYERS: int = 6 DROPOUT: float = 0.1 TASK_DIM: int = 512 # model components class HierarchicalAttention(nn.Module): def __init__(self, dim, heads=8): super().__init__() self.local_attn = nn.MultiheadAttention(dim, heads, batch_first=True) self.global_attn = nn.MultiheadAttention(dim, heads, batch_first=True) self.merge = nn.Linear(dim * 2, dim) self.task_gate = nn.Sequential( nn.Linear(dim, dim), nn.Sigmoid() ) def forward(self, x, task_embed=None): local_out = self.local_attn(x, x, x)[0] if task_embed is not None: x = x * self.task_gate(task_embed).unsqueeze(1) global_out = self.global_attn(x, x, x)[0] return self.merge(torch.cat([local_out, global_out], dim=-1)) class TransformerBlock(nn.Module): def __init__(self, config): super().__init__() self.norm1 = nn.LayerNorm(config.EMBED_DIM) self.attn = nn.MultiheadAttention( config.EMBED_DIM, config.NUM_HEADS, dropout=config.DROPOUT, batch_first=True ) self.norm2 = nn.LayerNorm(config.EMBED_DIM) self.mlp = nn.Sequential( nn.Linear(config.EMBED_DIM, config.EMBED_DIM * 4), nn.GELU(), nn.Dropout(config.DROPOUT), nn.Linear(config.EMBED_DIM * 4, config.EMBED_DIM) ) self.task_gate = nn.Sequential( nn.Linear(config.EMBED_DIM, config.EMBED_DIM), nn.Sigmoid() ) def forward(self, x, task): h = self.norm1(x) h = self.attn(h, h, h)[0] g = self.task_gate(task).unsqueeze(1) x = x + h * g h = self.norm2(x) h = self.mlp(h) x = x + h * g return x class WaveletTemporal(nn.Module): def __init__(self, config): super().__init__() self.embed_dim = config.EMBED_DIM self.spatial_proj = nn.Conv3d(1, config.EMBED_DIM, 1) self.temporal_proj = nn.Conv3d( config.EMBED_DIM, config.EMBED_DIM, (3,1,1), padding=(1,0,0) ) self.pool = nn.AdaptiveAvgPool3d((15, 32, 32)) def forward(self, x): b, t, h, d, w = x.shape x = x.reshape(b, 1, t, h, w*d) x = self.spatial_proj(x) x = self.temporal_proj(x) return self.pool(x) class SequentialBrainViT(nn.Module): def __init__(self, config): super().__init__() self.config = config self.temporal = WaveletTemporal(config) self.pool = nn.Sequential( nn.LayerNorm([config.EMBED_DIM, 15, 32, 32]), nn.AdaptiveAvgPool3d((5, 16, 16)), Rearrange('b c t h w -> b (t h w) c') ) self.num_patches = 5 * 16 * 16 self.task_embed = nn.Embedding(4, config.TASK_DIM) self.task_proj = nn.Sequential( nn.Linear(config.TASK_DIM, config.EMBED_DIM), nn.LayerNorm(config.EMBED_DIM), nn.GELU() ) self.cls_token = nn.Parameter(torch.zeros(1, 1, config.EMBED_DIM)) self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches + 1, config.EMBED_DIM)) self.blocks = nn.ModuleList([ TransformerBlock(config) for _ in range(config.NUM_LAYERS) ]) self.shared_proj = nn.Sequential( nn.LayerNorm(config.EMBED_DIM), nn.Linear(config.EMBED_DIM, config.EMBED_DIM * 2), nn.GELU(), nn.Linear(config.EMBED_DIM * 2, config.EMBED_DIM), nn.LayerNorm(config.EMBED_DIM), nn.Dropout(config.DROPOUT) ) self.heads = nn.ModuleDict({ 'learning_stage': nn.Sequential( nn.LayerNorm(config.EMBED_DIM), nn.Linear(config.EMBED_DIM, 1), nn.Sigmoid() ), 'region_activation': nn.Sequential( nn.LayerNorm(config.EMBED_DIM), nn.Linear(config.EMBED_DIM, 116) ), 'temporal_pattern': nn.Sequential( nn.LayerNorm(config.EMBED_DIM), nn.Linear(config.EMBED_DIM, 30) ) }) self._init_weights() def _init_weights(self): nn.init.normal_(self.cls_token, std=0.02) nn.init.normal_(self.pos_embed, std=0.02) for n, m in self.named_modules(): if isinstance(m, nn.Linear): nn.init.normal_(m.weight, std=0.02) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.LayerNorm): nn.init.ones_(m.weight) nn.init.zeros_(m.bias) def forward(self, x, task_ids): x = self.temporal(x) x = self.pool(x) cls_tokens = self.cls_token.expand(x.shape[0], -1, -1) x = torch.cat([cls_tokens, x], dim=1) x = x + self.pos_embed[:,:x.shape[1]] task = self.task_proj(self.task_embed(task_ids)) for block in self.blocks: x = block(x, task) x = self.shared_proj(x) return { 'learning_stage': self.heads['learning_stage'](x[:,0]), 'region_activation': self.heads['region_activation'](x.mean(1)), 'temporal_pattern': self.heads['temporal_pattern'](x[:,0]) } def preprocess_volume(vol, target_size=(64, 64, 30)): if vol.ndim == 4: vol = vol[None] b,t,h,w,d = vol.shape target_h, target_w, target_d = target_size vol = zoom(vol, ( 1, 1, target_h/h, target_w/w, target_d/d ), order=1) vol = (vol - vol.mean((1,2,3,4), keepdims=True)) / (vol.std((1,2,3,4), keepdims=True) + 1e-8) return torch.from_numpy(vol).float() def plot_results(region_acts, temporal_pattern): fig = plt.figure(figsize=(12,4)) plt.subplot(121) sns.heatmap(region_acts.reshape(1,-1), cmap='RdBu_r', center=0) plt.title('region activations') plt.xlabel('brain region') plt.subplot(122) plt.plot(temporal_pattern.squeeze()) plt.title('temporal pattern') plt.xlabel('time') return fig def process_fmri(file_obj): try: img = nib.load(file_obj.name) data = img.get_fdata(dtype=np.float32) if data.ndim != 4: return f"error: expected 4D data, got {data.ndim}D", None data = preprocess_volume(data) device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') results = {} figs = [] for stage in ['full', 'region', 'temporal']: model = SequentialBrainViT(Config()) ckpt = torch.load(f'best_{stage}.pt', map_location=device) model.load_state_dict(ckpt['model']) model.eval() with torch.no_grad(): outputs = model(data.to(device), torch.tensor([0]).to(device)) results[stage] = { 'learning_stage': float(outputs['learning_stage'].cpu().mean()), 'region_activation': outputs['region_activation'].cpu().numpy(), 'temporal_pattern': outputs['temporal_pattern'].cpu().numpy() } fig = plot_results( results[stage]['region_activation'], results[stage]['temporal_pattern'] ) figs.append(fig) plt.close() stage_results = "\n".join([ f"{stage.upper()} MODEL:" f"\nlearning stage: {res['learning_stage']:.3f}" f"\n" for stage, res in results.items() ]) return stage_results, figs[0] # return first fig for display except Exception as e: return f"error processing file: {str(e)}", None # create interface iface = gr.Interface( fn=process_fmri, inputs=gr.File(label="upload 4D fMRI nifti (.nii/.nii.gz)"), outputs=[ gr.Textbox(label="classification results"), gr.Plot(label="visualization") ], title="fmri learning stage classifier", description="upload a 4D fMRI nifti file to classify learning stages and visualize brain patterns", examples=[], cache_examples=False ) if __name__ == "__main__": iface.launch()