learnedSpectrum / app.py
twarner's picture
init
0ffa42f
Raw
History Blame
9.13 kB
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()