End-to-end fMRI-to-caption model using DynaDiff pretrained brain encoder and Transformer decoder.
Direct training pipeline for generating captions from fMRI brain signals:
- Brain Encoder: DynaDiff pretrained FmriMLP (346M parameters) - encodes fMRI signals to CLIP space embeddings
- Caption Decoder: Transformer decoder (trained from scratch) - generates captions from brain embeddings
- Training: Caption loss only (cross-entropy on next-token prediction)
- Evaluation: Comprehensive metrics on both captions and generated images
fMRI [B, voxels, TRs]
↓
DynaDiff Pretrained Brain Encoder (FmriMLP)
↓
Brain Embeddings [B, 257, 768] in CLIP Space
↓
Caption Decoder (Transformer)
↓
Generated Caption
↓
Cross-Entropy Loss
Key Design Choices:
- Leverages CLIP multimodal space for semantic alignment between brain signals and language
- Fine-tunes pretrained brain encoder while training decoder from scratch
- Direct optimization for caption quality (no image generation during training)
- Images generated post-hoc via Stable Diffusion for evaluation only
Ground truth captions are Qwen2-VL generated captions from original NSD images. We use high-quality vision-language model captions as ground truth for training and evaluation.
- Config:
config/train_config.yaml - Batch size: 8 (effective 64 with gradient accumulation)
- Learning rate: 1e-4
- Mixed precision: Enabled
- Gradient checkpointing: Enabled
sbatch slurm/train_caption.sbatch- BLEU-1, BLEU-2, BLEU-3, BLEU-4
- METEOR
- ROUGE-L
- CIDEr
Images generated from captions using Stable Diffusion, then evaluated:
- SSIM, PixCorr
- AlexNet(2), AlexNet(5)
- CLIP-12
- Inception V3
- EfficientNet, SWAV, DreamSim, mIoU
sbatch slurm/evaluate_caption_with_images.sbatch- Image-disjoint train/val/test splits (no COCO image overlap)
- Variable voxel count handling per subject
- Qwen2-VL tokenizer integration
- Ground truth: Qwen2-VL captions from original NSD images
model/brain_caption_model.py- Full model assemblymodel/brain_encoder.py- DynaDiff pretrained brain encodermodel/transformer_captioner.py- Caption decodertrain/train_caption_model.py- Training scripttest/evaluate_caption_with_images.py- Evaluation scriptdataPrepartion/fmri_caption_dataset.py- Dataset with image-disjoint splits
See INTERMEDIATE_RESULTS_SUMMARY.md for current evaluation results.
- PyTorch 2.4.0+
- transformers (Qwen2-VL tokenizer)
- diffusers (Stable Diffusion for evaluation)
- pycocoevalcap (caption metrics)
See requirements.txt for complete list.