Downloads ยท 30 days
0
praiselab-picuslab/BrainGemma3D
BrainGemma3D is a machine learning model from praiselab-picuslab. Use it for the machine learning task on the model card, and read the license before you ship it in a product. The card lists the license as cc-by-4.0.
BrainGemma3D is a multimodal vision-language model that generates clinically accurate radiology reports directly from native 3D brain MRI volumes. Unlike 2D slice-based approaches, BrainGemma3D processes MRI scans volโฆ
Downloads ยท 30 days
0
Access
Public
Updated Feb 24, 2026
Repo size
12.2 GB
Likes
3
Public
Click a slice to open those files.
.safetensors12.1 GB ยท 100%
From the Hugging Face model README
BrainGemma3D is a multimodal vision-language model that generates clinically accurate radiology reports directly from native 3D brain MRI volumes. Unlike 2D slice-based approaches, BrainGemma3D processes MRI scans volumetrically, preserving the spatial context critical for accurate neuroradiological interpretation.
<div align="center"> <a href="https://github.com/PRAISELab-PicusLab/BrainGemma3D" target="_blank"><img alt="GitHub Repository" src="https://img.shields.io/badge/GitHub-BrainGemma3D-181717?style=for-the-badge&logo=github&logoSize=auto"/></a> <a href="https://www.kaggle.com/code/antonioromano45/braingemma3d" target="_blank"><img alt="Kaggle Notebook" src="https://img.shields.io/badge/Kaggle-Notebook-20BEFF?style=for-the-badge&logo=kaggle&logoSize=auto"/></a> <br> <a href="https://www.kaggle.com/competitions/med-gemma-impact-challenge/overview" target="_blank"><img alt="MedGemma Challenge" src="https://img.shields.io/badge/Kaggle-MedGemma_Impact_Challenge-blue?style=for-the-badge&logo=kaggle&logoSize=auto&color=20BEFF"/></a> </div>BrainGemma3D combines:
3D Vision Encoder: MedSigLIP inflated to 3D via center-frame initialization (Conv2D โ Conv3D)
Base model: google/medsiglip-448
Token Compressor: 2-layer Perceiver that reduces 3D patches to 32 visual tokens
Vision-Language Projector: 2-layer MLP that projects visual tokens to language model embedding space
Language Model: 4-bit quantized MedGemma-1.5-4B-IT with LoRA adapters
Base model: google/medgemma-1.5-4b-it
pip install torch torchvision transformers nibabel scikit-image lime
from huggingface_hub import snapshot_download
# 1. Download the repository containing our custom architecture from Hugging Face
repo_id = "praiselab-picuslab/BrainGemma3D"
print(f"Downloading repository: {repo_id}...")
local_dir = snapshot_download(repo_id)
print(f"โ
Repository downloaded to: {local_dir}")
import os
import torch
import sys
sys.path.append(local_dir)
from medgemma3d_architecture import MedGemma3D, load_nifti_volume, CANONICAL_PROMPT
# Automatically select the optimal hardware accelerator (GPU if available, otherwise CPU)
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Hardware accelerator selected: {device}")
# 2. Instantiate the base architecture (3D-inflated MedSigLIP + MedGemma)
model = MedGemma3D(
vision_model_dir=f"{local_dir}/vision_model",
language_model_dir=f"{local_dir}/language_model",
depth=2,
num_vision_tokens=32,
freeze_vision=True,
freeze_language=True,
device_map={"": 0} if device == "cuda" else None,
)
# 3. Load projector
proj_path = os.path.join(local_dir, "projector_vis_scale.pt")
print(f"Loading custom projector weights from: {proj_path}...")
# Load the checkpoint into memory
ckpt = torch.load(proj_path, map_location=device)
# Inject the weights into the visual projector (which bridges Vision and Language)
model.vision_projector.load_state_dict(ckpt["vision_projector"])
# Load the visual scaling factor, ensuring correct tensor formatting
if ckpt.get("vis_scale") is not None:
if isinstance(ckpt["vis_scale"], torch.Tensor):
model.vis_scale.data = ckpt["vis_scale"].to(device)
else:
model.vis_scale.data.fill_(ckpt["vis_scale"])
# Transition the model to evaluation mode for inference
model.eval()
print("โ
BrainGemma3D is fully loaded and ready for inference!")
# 4. Load MRI scan
volume = load_nifti_volume(
"path/to/brain_flair.nii.gz",
target_size=(32, 128, 128)
).to(device)
if volume.ndim == 4:
volume = volume.unsqueeze(0)
# 5. Generate report
with torch.no_grad():
report = model.generate_report(
volume,
prompt=CANONICAL_PROMPT,
max_new_tokens=256,
temperature=0.1,
top_p=0.9,
)
print("\n===== GENERATED REPORT =====\n")
print(report)
Generated Report:
The lesion area is in the left parietal and frontal lobes with mixed high-signal
areas. Edema signals are mainly observed around these lesions, indicating significant
edema presence affecting parts of both frontal and temporal regions as well as some
portions within the parietal lobe. Necrosis may be present at low signal intensity
or scattered throughout certain sections of the brain tissue affected by edema.
Ventricular compression effects on adjacent ventricles can occur due to pressure
from surrounding tissues near the ventricular system.
BrainGemma3D is trained in three progressive stages to prevent catastrophic forgetting:
Dataset:
Evaluated on 468 subjects (369 BraTS pathological + 99 healthy controls) with group-based splits.
| Model | BLEU-1 | BLEU-4 | ROUGE-L | CIDEr | Lat F1 | Anat F1 | Path F1 |
|---|---|---|---|---|---|---|---|
| Med3DVLM (3D Generalist) | 0.051 | 0.005 | 0.083 | 0.007 | 0.300 | 0.225 | 0.119 |
| MedGemma 1.5 (2D Slice) | 0.245 | 0.024 | 0.189 | 0.029 | 0.526 | 0.461 | 0.413 |
| BrainGemma3D (Ours) | 0.302 | 0.098 | 0.289 | 0.293 | 0.689 | 0.691 | 0.951 |
Key Insight: The +130% gain in Pathology F1 (0.951 vs 0.413 compared to 2D baseline) demonstrates that native 3D processing is essential for diagnostic accuracy in neuroradiology.
BrainGemma3D includes LIME-based 3D interpretability to visualize which brain regions drive diagnostic predictions.
from braingemma3d_interpretability import run_interpretability
# 6. Run interpretability analysis
weights, wvol = run_interpretability(
model=model,
load_nifti_volume=load_nifti_volume,
CANONICAL_PROMPT=CANONICAL_PROMPT,
mri_path="path/to/brain_flair.nii.gz",
report=report,
output_dir="./interpretability_output",
lime_samples=100, # Number of perturbations (more = better but slower)
n_segments=20, # Number of brain regions to analyze
alpha=0.45, # Overlay transparency
clip_q=0.99, # Heatmap clipping
seed=42,
)
Output:
overlay_slices.png โ Full 3D heatmap (red=supportive, blue=contradicting)lime_2x3_grid.png โ 2ร3 grid with selected slices (original + LIME overlay)lime_top_supervoxels_grid.png โ Most influential supervoxelslime_weights.json โ Supervoxel importance scoresBrainGemma3D achieved 95.1% pathology F1 on the BraTS, but this does NOT imply clinical readiness. Key considerations:
Recommendation: Use only in research settings with appropriate ethical oversight and informed consent.
This project was developed by:
Mariano Barone ยท Francesco Di Serio ยท Giuseppe Riccio ยท Antonio Romano ยท Vincenzo Moscato
Department of Electrical Engineering and Information Technology
University of Naples Federico II, Italy