Skip to content

Repository files navigation

MedSAM with CLIP Integration for Medical Image Segmentation

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.

Overview

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.

Dataset

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.

Model Architecture

Baseline MedSAM

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.

Enhanced MedSAM with CLIP

Our enhanced version integrates CLIP models for improved text-image understanding and cross-modal attention mechanisms, enabling better organ segmentation through text prompts.

Installation

Prerequisites

  • Python 3.9+
  • CUDA-compatible GPU (recommended)

Setup

  1. Clone the repository:
git clone https://github.com/your-username/Prompt-Dimensions-of-MedSAM.git
cd Prompt-Dimensions-of-MedSAM
  1. Install dependencies:
pip install -r requirements.txt
  1. 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

Data Preparation

FLARE 2022 Dataset

  1. 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
  1. Preprocess the data using the provided utilities in utils/
  2. 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

Data Preprocessing

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)

Training

Full Supervised Training

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"

Key Parameters

  • --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

Model Architecture Details

CLIP Integration

The model supports multiple CLIP variants:

  • BiomedCLIP: Medical domain-specific CLIP model
  • Standard CLIP: General-purpose CLIP model
  • BioCLIP: Biology-focused CLIP model

Cross-Modal Attention

Implements multi-head cross-attention between image and text features for enhanced segmentation performance.

Supported Organs

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

Evaluation

The model outputs include:

  • Training and validation loss curves
  • Dice coefficient scores for each organ
  • Model checkpoints (latest and best)
  • Training visualization plots

File Structure

├── 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

Dependencies

  • 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)

Citation

If you use this code in your research, please cite:

License

This project is licensed under the MIT License - see the LICENSE file for details.

Acknowledgments

  • Original SAM implementation from Meta AI
  • FLARE 2022 challenge organizers
  • BiomedCLIP and related CLIP variants

About

MedSAM enhanced with CLIP text descriptors for multi-organ medical image segmentation, evaluated on 24k labeled CT slices from FLARE 2022.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages