Instruction-conditioned diffusion medical image segmentation using Stable Diffusion v1.5 carried by a frozen DINOv3 visual conditioner.
Users provide a natural-language instruction (e.g. "segment the polyp region using red") and the model generates the target label directly as a color-coded overlay on the original image, in a single forward diffusion process without any task-specific prediction head.
[train.py] -> [inference.py]
│ │
├─ dataset/ ├─ models/sd15/ (SDv1.5 VAE + UNet)
├─ editor_utils.py ├─ models/dinov3_pyramid.py (multi-scale visual feature pyramid)
└─ (BTCV/2D data) ├─ models/aux_backbones.py (ablation backbones)
└─ models/dinov3/ (official DINOv3 library)
- Python 3.10
- PyTorch 2.x,
diffusers,transformers segment-anything(only for the SAM auxiliary-backbone ablation, when it is used)timm(transitive dependency of the bundled DINOv3 library)
pip install torch diffusers transformers segment-anything timm tqdm pillowDownload and place into assets/ (paths are resolved relative to the project root):
| Component | Location |
|---|---|
| Stable Diffusion v1.5 | https://huggingface.co/runwayml/stable-diffusion-v1-5 → assets/sd-v1-5/ |
| DINOv3 ViT-S/16 | assets/dinov3_vits16_pretrain_lvd1689m-08c60483.pth |
| DINOv3 ViT-B/16 | assets/dinov3_vitb16_pretrain_lvd1689m-73cec8be.pth |
| DINOv3 ViT-L/16 | assets/dinov3_vitl16_pretrain_lvd1689m-8aa4cbdd.pth |
| OpenAI CLIP ViT-B/16 | assets/ablation_backbones/clip_vitb16/ (CLIPVisionModel weights) |
| DINOv2 ViT-B/14 | assets/ablation_backbones/dinov2_vitb14/ |
| ImageNet ResNet-50 | assets/ablation_backbones/resnet50/ |
| MedSAM ViT-B image encoder | assets/medsam_vit_b.pth |
DINOv3 weights are released by the official DINOv3 repository (the dinov3 library is vendored in models/dinov3, with its original license in models/dinov3/LICENSE.md).
Ablation backbones (CLIP/DINOv2/ResNet/SAM) can be copied from a HuggingFace mirror, e.g. via huggingface-cli download --local-dir.
Each dataset is a JSONL file of samples, stored at data/Instruction_{dataset}/train.jsonl and test.jsonl. Each line:
{
"image": "path/to/sample.png",
"label": "path/to/mask.png",
"class_indices": [0, 1],
"palette": {"0": "#000000", "1": "#ff0000"}
}The dataset classes map this to per-sample fields pixel_values, labels (overlay), original_mask, color_map, class_indexes and a generated instruction.
dataset/prepare_polypgen.py– build the polyp_combined instruction set (Kvasir/CVC/ETIS/PolypGen).dataset/build_sliced_jsonl.py– build BTCV axial-slice instruction set.
Main protocol (paper): batch size 8, seed 42, 20k steps, AdamW lr 1e-4 (UNet/DINOv3Pyr) / 5e-7 (CLIP text encoder), DINOv3 ViT-S/16 with the last 6 transformer blocks unfrozen, auxiliary segmentation loss weight λ = 0.3, overlay target.
python train.py \
--dataset polyp_combined --use_dinov3 \
--train_metadata data/Instruction_polyp_combined/train.jsonl \
--val_metadata data/Instruction_polyp_combined/test.jsonl \
--output_dir output/sd15_polyp_combined_v6 \
--max_train_steps 20000 --batch_size 8 --seed 42 --device cuda:0 \
--seg_loss_weight 0.3 --dinov3_seg_loss_weight 0.3 # cf. paperBTCV multi-organ CT follows the same pipeline, adding the 14-class palette and its own metadata:
python train.py --dataset btcv --use_dinov3 --max_train_steps 20000 --device cuda:0# Output representation: binary grayscale mask / color mask on black background
python train.py --dataset polyp_combined --use_dinov3 --target_mode mask_binary --output_dir output/ablation/polyp/target_mask_binary --max_train_steps 10000 --seed 42
python train.py --dataset polyp_combined --use_dinov3 --target_mode mask_black_bg ...
# Auxiliary backbones
python train.py --dataset polyp_combined --no_dinov3 ... # None
python train.py --dataset polyp_combined --use_dinov3 --backbone_type clip_vitb16 ...
python train.py --dataset polyp_combined --use_dinov3 --backbone_type dinov2_vitb14 ...
python train.py --dataset polyp_combined --use_dinov3 --backbone_type resnet50 ...
python train.py --dataset polyp_combined --use_dinov3 --backbone_type sam_vitb ...
python train.py --dataset isic --use_dinov3 --dinov3_model_size vitl16 \
--dinov3_weights_path assets/dinov3_vitl16_pretrain_lvd1689m-8aa4cbdd.pth ...See scripts/run_ablation_train.sh for the full matrix.
python inference.py \
--checkpoint_path output/sd15_polyp_combined_v6_small_segloss/final \
--dataset kvasir --metadata_path data/Instruction_kvasir/test.jsonl \
--output_dir output/inference/kvasir --use_instruction --device cuda:0Single sample: add --index 3. Supported datasets: kvasir, isic, isic2017, cvc_clinicdb, cvc_colondb, etis, polypgen (2D) and btcv (axial CT).
Notes:
- The paper's unseen-color experiments prompt with
--force_color(e.g.--force_color red) and restrict to a subset with--index;BINARY_COLORS_V6may be extended indataset/instruction_2d_dataset_v6.pyfor colors outside the 16-color pool. --backbone_type: override the spatial backbone. The output-representation ablations (mask_binary,mask_black_bg) were trained/evaluated with the modern DINOv3 wrapper — rerun with--backbone_type vits16for exact table reproduction; the main model checkpoint uses the classic DINOv3 path (default).
Reproducibility: DDIM, 25 steps, text CFG 7.5, noise level 1.0, seed 42. The evaluation computes Dice/IoU per foreground class with the fixed color map, plus a small-component filter (>= 18 px) for the binary/overlay mask extraction.
train.py training entry (UNet + CLIP + DINOv3Pyramid)
inference.py evaluation entry (single sample / full dataset + metrics CSVs)
editor_utils.py shared pure helpers
dataset/ datasets, losses, metrics, data preparation
models/sd15 SD v1.5 VAE/UNet
models/dinov3_pyramid.py multi-scale feature pyramid + auxiliary seg head
models/aux_backbones.py ablation backbones (CLIP/DINOv2/ResNet-50/MedSAM)
models/dinov3/ official DINOv3 library (Meta, see LICENSE)
scripts/ main / ablation train + eval driver scripts
MIT (see LICENSE). The vendored DINOv3 library is distributed under its own license (models/dinov3/LICENSE.md); the MedSAM image encoder weights are obtained from the official MedSAM release.