This repository contains the implementation of MedSAM (Medical Segment Anything Model) enhanced with CLIP (Contrastive Language-Image Pre-training) integration for multi-organ medical image segmentation on the FLARE 2022 dataset.
We adopt the original MedSAM as the base model and load its official weights. During training and testing, the model is trained for 30 epochs on the FLARE 2022 dataset with a batch size of 8, using the Adam optimizer (learning rate: 0.0001). The baseline models are also trained for 30 epochs on the FLARE 2022 dataset with the same hyperparameters.
All experiments were conducted on the FLARE 2022 dataset, a benchmark comprising 40 contrast-enhanced abdominal CT volumes (cases 0001-0040) with expert annotations for 13 organs: liver, spleen, pancreas, stomach, gallbladder, duodenum, esophagus, aorta, inferior vena cava (IVC), left and right kidneys, and left and right adrenal glands. Each volume is provided in 3D NIfTI format.
To adapt this data to our 2D segmentation pipeline, we resampled scans to a uniform pixel spacing and stored them as compressed NumPy archives (.npz), resulting in approximately 24,000 fully labeled slices. For consistency across organs, any class with fewer than 2,000 labeled slices was excluded; consequently, the duodenum was omitted, leaving 12 organs for training, validation, and testing.
MedSAM processes 2D slices using a frozen ViT-B/16 image encoder pretrained on large-scale natural and medical image datasets. The encoder converts input slices into patch-level visual tokens. A separate prompt encoder embeds user-provided guidance (e.g., points or bounding boxes) into tokens. These visual and prompt tokens are concatenated and passed to a mask decoder, which outputs a probability map for the target structure.
Our enhanced version integrates CLIP models for improved text-image understanding and cross-modal attention mechanisms, enabling better organ segmentation through text prompts.
- Python 3.9+
- CUDA-compatible GPU (recommended)
- Clone the repository:
git clone https://github.com/your-username/Prompt-Dimensions-of-MedSAM.git
cd Prompt-Dimensions-of-MedSAM- Install dependencies:
pip install -r requirements.txt- Download the SAM checkpoint:
# Create the SAM directory
mkdir -p work_dir/SAM
# Download sam_vit_b_01ec64.pth from the official SAM repository
wget https://dl.fbaipublicfiles.com/segment_anything/sam_vit_b_01ec64.pth -O work_dir/SAM/sam_vit_b_01ec64.pth
# Alternative download link: https://github.com/facebookresearch/segment-anything/releases/download/v1.0/sam_vit_b_01ec64.pth
# Verify the file is placed correctly
ls -la work_dir/SAM/sam_vit_b_01ec64.pth
# Note: This file is ~358MB and not included in the repository due to GitHub file size limits
# The file must be placed in work_dir/SAM/ directory for the training script to work properly- Download the original FLARE 2022 dataset:
# Download from the official FLARE 2022 challenge website
# https://flare.grand-challenge.org/
# Or download from Zenodo:
# https://zenodo.org/records/5903672
# Note: You need to download the original dataset and preprocess it yourself
# The original dataset contains 50 cases, we use cases 0001-0040 for training- Preprocess the data using the provided utilities in
utils/ - Organize the data in the following structure:
data/
└── npy/
└── CT_Abd/ # Preprocessed dataset directory
├── CT_Abd_FLARE22_Tr_0001.npz # Case 0001 (image + ground truth)
├── CT_Abd_FLARE22_Tr_0002.npz # Case 0002
├── ... # Cases 0003-0039
├── CT_Abd_FLARE22_Tr_0040.npz # Case 0040
├── CT_Abd_FLARE22_Tr_0001_img.nii.gz # Original CT image (optional)
├── CT_Abd_FLARE22_Tr_0001_gt.nii.gz # Original annotation (optional)
└── ... # Additional original files
Dataset Structure:
- 40 training cases (0001-0040) from FLARE 2022 challenge
- Each case contains: CT image slices and corresponding organ annotations
- File formats:
.npz: Preprocessed NumPy arrays (recommended for training)_img.nii.gz: Original CT images in NIfTI format_gt.nii.gz: Original annotations in NIfTI format
After downloading the FLARE 2022 dataset, use the provided preprocessing scripts in the utils/ directory:
# Convert CT scans to 2D slices and organize by organs
python utils/pre_CT_MR.py # CT/MR preprocessing
# Convert grayscale images to RGB format
python utils/pre_grey_rgb.py # Grayscale to RGB conversion
# Split data into training/validation sets
python utils/split.py # Data splitting utilities
# Note: The preprocessing will create the required directory structure
# and convert the 3D NIfTI files to 2D NumPy arrays (.npz format)Run the complete supervised training script:
python train_medsam_full_supervised.py \
--tr_npy_path data/npy/CT_Abd \
--task_name "MedSAM-ViT-B" \
--model_type "vit_b" \
--checkpoint "./work_dir/SAM/sam_vit_b_01ec64.pth" \
--work_dir "./work_dir1" \
--num_epochs 30 \
--batch_size 8 \
--lr 0.0001 \
--use_clip True \
--clip_variant "biomedclip" \
--device "cuda:0"--use_clip: Enable CLIP integration (default: True)--clip_variant: CLIP model variant ("biomedclip", "clip", "bioclip")--ms_features: Enable multi-scale features--one_neck: Use single neck architecture--use_amp: Enable automatic mixed precision training--use_wandb: Enable Weights & Biases logging
The model supports multiple CLIP variants:
- BiomedCLIP: Medical domain-specific CLIP model
- Standard CLIP: General-purpose CLIP model
- BioCLIP: Biology-focused CLIP model
Implements multi-head cross-attention between image and text features for enhanced segmentation performance.
The model supports segmentation of 12 organs:
- Liver
- Right kidney
- Spleen
- Pancreas
- Aorta
- Inferior vena cava (IVC)
- Right adrenal gland
- Left adrenal gland
- Gallbladder
- Esophagus
- Stomach
- Left kidney
The model outputs include:
- Training and validation loss curves
- Dice coefficient scores for each organ
- Model checkpoints (latest and best)
- Training visualization plots
├── train_medsam_full_supervised.py # Main training script
├── test_data.py # Dataset loading and preprocessing
├── get_clip_embedding1.py # CLIP integration module
├── requirements.txt # Python dependencies
├── README.md # This file
├── data/ # Dataset directory
│ └── npy/
├── segment_anything/ # SAM model implementation
├── utils/ # Utility functions
│ ├── SurfaceDice.py # Dice coefficient computation
│ ├── pre_CT_MR.py # CT/MR preprocessing
│ ├── pre_grey_rgb.py # Image format conversion
│ └── split.py # Data splitting utilities
└── work_dir/ # Model checkpoints and outputs
└── SAM/ # SAM model weights
- PyTorch >= 1.9.0
- MONAI >= 1.0.0
- Transformers >= 4.20.0
- OpenCLIP >= 2.0.0
- Torchvision >= 0.10.0
- Scikit-image >= 0.19.0
- Nibabel (for medical image processing)
- Weights & Biases (optional, for experiment tracking)
If you use this code in your research, please cite:
This project is licensed under the MIT License - see the LICENSE file for details.
- Original SAM implementation from Meta AI
- FLARE 2022 challenge organizers
- BiomedCLIP and related CLIP variants