Skip to content

Latest commit

 

History

2 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

InstEditSeg

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)

Requirements

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

Pretrained assets

Download 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.

Data format

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.

Training

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. paper

BTCV 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

Ablation variants

# 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.

Inference / Evaluation

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:0

Single 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_V6 may be extended in dataset/instruction_2d_dataset_v6.py for 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 vits16 for 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.

Repository layout

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

License

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.

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages