Supervised fine-tuning (SFT)#
SFT teaches a vision-language model to produce MedVision’s structured measurements from a medical image plus an instruction. The reference recipes fine-tune Qwen2.5-VL-7B-Instruct with chain-of-thought (CoT) targets: the response reasons through the measurement (landmark coordinates, pixel geometry) before emitting the final value, so the model learns the procedure rather than memorising numbers.
All training runs through python -m medvision_bm.sft.<entry-point> argparse drivers. You rarely call them by hand — the shell scripts under script/sft/ wire up the environment, the two-phase pipeline, and the full flag list for you.
Note
This page assumes the package and data are already in place. See Installation for the environment and the MedVision_* variables, and Dataset loading for how task-list JSONs resolve to samples.
The three recipes#
script/sft/ ships three ready-to-run scripts, all training the same 121K-sample multi-task mix (110K Detection + 5.5K Angle/Distance + 5.5K Tumour/Lesion):
Script |
Method |
Resolution |
Launcher |
|---|---|---|---|
|
LoRA adapters |
native / dynamic |
DDP |
|
LoRA adapters |
512×512 |
DDP |
|
full-parameter |
512×512 |
FSDP |
The LoRA scripts train adapters on a frozen backbone and launch with plain DistributedDataParallel. The full-parameter script updates every weight; at 7B that does not fit in DDP on 80 GB GPUs (weights + gradients + FP32 AdamW state ≈ 84 GB/GPU before activations), so it shards optimizer state, gradients, and parameters across GPUs with FSDP.
The __512x512 variants add --new_shape_hw 512 512, which resizes each slice during dataset preparation and re-derives the physical pixel size for that resolution. Because measurement tasks depend on knowing the real millimetre-per-pixel scale, the prompt’s pixel size always matches the resolution the model actually perceives — the 512×512 full SFT recipe is the one behind the released MedVision-V0 checkpoints.
Note
MedVision-V0 is produced by two-stage post-training: this full-parameter 512×512 SFT, followed by reinforcement fine-tuning (GRPO). See Reinforcement fine-tuning.
Beyond these 7B reference recipes, script/sft/ carries the same layout for larger families — MedGemma-27B, Gemma-4-31B, Qwen3.5/Qwen3.6-27B — whose memory-recipe variants are covered in Scaling full-parameter SFT to 27B and beyond below.
To run one, set the paths and identifiers at the top of the script (benchmark_dir, data_dir, base_model_hf, run_name, W&B fields, and the batch/GPU settings) and execute it from the repo root:
bash script/sft/train__SFT-CoT__Qwen2.5VL7B__D110k-AD5.5k-TL5.5k__512x512.sh
Each script first provisions a dedicated conda env (sft-qwen25vl), builds medvision_bm into a wheel, and installs the model-specific extras:
python -m medvision_bm.sft.env_setup --data_dir ${data_dir} --lmms_eval_opt_deps qwen2_5_vl
The scripts pin the planner version and acknowledge the release (required whenever you pin below latest):
export MedVision_PLANNER_VERSION='1.0.0'
export MedVision_ACK_RELEASE='1.1.1'
Two-phase pipeline: prepare on CPU, train on GPU#
Building the prepared dataset for 121K samples — slicing NIfTI volumes, normalising, formatting CoT targets, caching PNGs — is CPU-bound and slow enough to trip distributed-training timeouts if done inside the training job. So every script invokes the same entry module twice.
Phase 1 — dataset preparation (CPU, single process). Runs with --process_dataset_only true, which downloads and formats every sample and writes the prepared dataset to disk. --save_processed_img_to_disk true also emits processed slices as PNGs so training loads them directly instead of re-slicing volumes:
python -m medvision_bm.sft.train__SFT-CoT__qwen2_5_vl \
--process_dataset_only true \
--skip_process_dataset false \
--save_processed_img_to_disk true \
--data_dir ${data_dir} \
--model_family_name qwen25vl \
--base_model_hf Qwen/Qwen2.5-VL-7B-Instruct \
--new_shape_hw 512 512 \
... # task lists + sample limits (see below)
The prepared dataset lands in --prepared_ds_dir, defaulting to a path derived from the per-task limits, e.g. <data_dir>/tmp_prepared_ds_AD5500_D110000_TL5500_all121000.
Phase 2 — training (GPU, distributed). Launched under accelerate with --skip_process_dataset true so it loads the cached dataset instead of rebuilding it.
LoRA uses a plain DDP launch:
CUDA_VISIBLE_DEVICES=0,1,2,3 \
accelerate launch --num_processes=4 --main_process_port=29502 --mixed_precision=bf16 \
-m medvision_bm.sft.train__SFT-CoT__qwen2_5_vl \
--skip_process_dataset true \
--process_dataset_only false \
...
Full-parameter training swaps in the train__fullFT-CoT__qwen2_5_vl module and adds FSDP flags to shard the model. Note the transformer layer class to wrap is model-specific (Qwen2_5_VLDecoderLayer for Qwen2.5-VL):
CUDA_VISIBLE_DEVICES=0,1,2,3 \
accelerate launch --num_processes=4 --main_process_port=29502 --mixed_precision=bf16 \
--use_fsdp \
--fsdp_sharding_strategy FULL_SHARD \
--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \
--fsdp_transformer_layer_cls_to_wrap Qwen2_5_VLDecoderLayer \
--fsdp_state_dict_type FULL_STATE_DICT \
--fsdp_offload_params false \
--fsdp_cpu_ram_efficient_loading true \
--fsdp_sync_module_states true \
-m medvision_bm.sft.train__fullFT-CoT__qwen2_5_vl \
--skip_process_dataset true \
...
Tip
Because phase 2 only reads the cache, re-running with --skip_process_dataset true skips preparation entirely. Combined with --resume_from_checkpoint true, an interrupted run under the same run_name simply picks up from its last checkpoint.
Multi-task inputs and sample limits#
Tasks enter training as task-list JSONs, one flag per task; supply at least one, or several for joint multi-task training:
--tasks_list_json_path_AD tasks_list/tasks_MedVision-AD__train_SFT.json
--tasks_list_json_path_detect tasks_list/tasks_MedVision-detect__train_SFT.json
--tasks_list_json_path_TL tasks_list/tasks_MedVision-TL__train_SFT.json
Global caps --train_sample_limit and --val_sample_limit are always required. On top of them you pick one of two balancing strategies:
Balanced —
--train_sample_limit_per_task/--val_sample_limit_per_taskspread the budget roughly evenly across the three tasks.Per-task (the shipped setting) —
--train_sample_limit_task_AD,--train_sample_limit_task_Detection,--train_sample_limit_task_TL(and their--val_...counterparts) set exact counts, e.g. 5.5K / 110K / 5.5K.
If a limit exceeds the available samples for a task, it is a no-op: the pool is capped at what is available and never oversampled or repeated. (The only with-replacement oversampling is the optional temperature sampler, enabled with --enable_temperature_sampler, which rebalances the multi-task mix by task frequency and is independent of these limits.)
Key hyperparameters#
These are the knobs the scripts expose most often; they map straight to SFTConfig/TrainingArguments:
Flag |
Role |
Recipe defaults |
|---|---|---|
|
training epochs |
|
|
per-GPU batch |
LoRA |
|
accumulation; effective batch = per-device × accum × #GPUs |
|
|
trade compute for memory |
|
|
FlashAttention-2 kernels |
|
|
resize + rescale pixel size in prep |
|
|
checkpoint / eval / log cadence |
|
|
max retained checkpoints |
|
|
resume the same |
|
The model family is chosen with --model_family_name (e.g. qwen25vl) plus --base_model_hf (a Hub ID or local path). The family name is validated at startup against the registered model list — both vllm_qwen25vl and the bare qwen25vl are accepted — so a typo fails fast instead of mid-run.
Temperature-based multi-task sampling#
With 110K detection samples against 5.5K each for A/D and T/L, uniform sampling lets detection swamp every batch. Turning on the temperature sampler re-weights how often each task is drawn:
--enable_temperature_sampler true \
--temperature_sampler_T 5
Internally this swaps the trainer for a TemperatureSamplerSFTTrainer subclass whose train sampler is a WeightedRandomSampler (with replacement, seeded from the project SEED). Per-task probability is count^(1/T), normalised, and each sample’s weight is that task probability divided by the task’s count. T = 1 reproduces count-proportional sampling; larger T flattens the distribution so the minority tasks are oversampled — the scripts use T = 5. It only reshapes training batches and has no effect during phase-1 preparation. With a single task present, it transparently falls back to the standard sampler.
Loss masking (completion-only)#
Which tokens count toward the training loss is decided per model family in the collate functions:
Qwen collates (shared by the Qwen2.5-VL and Qwen3-VL drivers) are completion-only by default: only each assistant response and its closing turn marker stay in the loss; padding, image tokens, and the entire user prompt are set to the
-100ignore index.Gemma-family collates (MedGemma, Gemma 4) mask only padding and image tokens by default, so the user prompt is part of the language-modeling loss — the same objective as Google’s official MedGemma fine-tuning notebook, whose collator masks exactly those tokens.
Setting MEDVISION_SFT_COMPLETION_ONLY=1 switches the Gemma-family collates to the same completion-only objective (the Qwen collates already mask, and ignore the flag). It applies to LoRA and full-parameter training alike, and raises at the first batch — rather than silently mis-masking — if a checkpoint’s chat template lacks the expected Gemma turn markers.
Warning
train/loss is not comparable across this flag: with masking on, the loss averages over only the response tokens (roughly 15 % of the sequence — all of them the hard answer tokens) instead of being diluted by near-identical prompt boilerplate, so the reported value jumps up. That is the flag working, not a regression.
And because it is an environment variable, a MEDVISION_SFT_COMPLETION_ONLY=1 left exported in the shell silently turns the next baseline Gemma run into a completion-only run. Launch the __cmplLoss script variants (which export this flag) in a fresh shell, or unset the variable before a baseline run; a sudden train/loss jump is the tell.
For MedVision’s long-CoT targets the expected downstream-accuracy effect is roughly neutral; the motivation is objective consistency with the Qwen family rather than an accuracy gain. The evidence review lives in the repository at docs/literature-review__loss-masking-in-SFT.md.
Scaling full-parameter SFT to 27B and beyond#
script/sft/ ships full-parameter CoT recipes for MedGemma-27B, Gemma-4-31B, and Qwen3.5/Qwen3.6-27B in two memory recipes per family:
Script variant |
Recipe |
Checkpoints |
Hardware |
|---|---|---|---|
|
anti-OOM: pure bf16 + 8-bit AdamW |
weights-only (~54 GB at 27B) |
4× 80 GB |
|
fp32 master weights + fused fp32 AdamW |
fully resumable (~160 GB at 27B) |
4× 140 GB-class |
__cmplLoss variants of either recipe additionally export the loss-masking flag above.
The standard AMP setup (bf16 compute, fp32 master weights, fused fp32 AdamW) carries ~121.5 GB of fixed per-GPU state at 27B across 4 ranks — far beyond 80 GB cards. The anti-OOM recipe trains bf16-native instead: MEDVISION_SFT_PURE_BF16=1 removes the fp32 masters (the launch also omits --mixed_precision=bf16), MEDVISION_SFT_OPTIM=adamw_bnb_8bit shrinks optimizer state 8×, and the fixed cost drops to ~40.5 GB. The learning rate is raised to 4e-5 so AdamW updates clear bf16’s ~0.4 % rounding resolution.
Warning
The anti-OOM recipe is not fully resumable. Its 8-bit optimizer state cannot be gathered by FSDP’s FULL_STATE_DICT, so checkpoints are weights-only (MEDVISION_SFT_SAVE_ONLY_MODEL=1): any restart — preemption, pod loss, a deliberate stop — discards the optimizer moments and LR-schedule position and warm-restarts from the last saved weights.
Prefer the fp32-master recipe whenever 140 GB-class GPUs are available; reserve the anti-OOM recipe for 80 GB pods. It keeps standard AMP numerics and fully resumable checkpoints, and fits 4 ranks — validated at 27B (Qwen3.6: post-FSDP-wrap ≈ 38 GiB/rank). At 31B its worst-case fixed cost (~139.5 GB/rank) sits at the edge of the budget: treat Gemma-4-31B on 4 GPUs as unvalidated and watch the memory probes on the first run.
The knobs are environment variables exported by the launcher scripts (not argparse flags). Every knob defaults to the legacy behavior, so the 7B pipelines are untouched:
Env var |
Default |
Effect when set |
|---|---|---|
|
off |
disable AMP ( |
|
|
any |
|
off |
weights-only checkpoints; required with 8-bit optimizers under FSDP |
|
|
full-FT learning-rate override |
|
off |
neutralise |
|
follows |
attention-implementation override (e.g. |
|
off |
Liger kernels; the fused cross-entropy removes the vocab-sized logits spike (requires |
|
off |
log per-rank memory after FSDP wrap and after step 1 |
|
off |
on OOM, dump a per-rank CUDA allocator snapshot into the checkpoint dir |
Resuming a full-parameter run is FSDP-aware: the entry points detect the last checkpoint before building the trainer and load its weights through the same sharded from_pretrained path as a fresh start, skipping the Trainer’s own checkpoint loader (which would all-gather the full unsharded model on every rank — an OOM at 27B+). This happens automatically with --resume_from_checkpoint true, and a checkpoint moves freely between pods with different GPU memory as long as the effective batch (world size × accumulation × per-device batch) stays the same.
Tip
MEDVISION_SFT_MEMPROBE=1 costs two log lines. On any new model-size/GPU combination, check the post-wrap figure first: expect roughly params × 2 bytes / world_size for pure bf16, or that plus the fp32 masters under AMP — a full-model-sized figure means FSDP sharding did not engage.
Background for these choices is collected in the repository at docs/literature-review__anti-OOM-fullFT-techniques.md.
Merging and pushing (LoRA only)#
The LoRA drivers can merge the trained adapter back into the base model and push either artifact to the Hub:
Flag |
Effect |
|---|---|
|
after training, merge the final adapter into the base weights |
|
local output dir and Hub repo name for the merged model |
|
upload the merged model to the Hub |
|
upload the LoRA adapter after each save |
|
skip training; merge and push the last existing checkpoint |
The full-parameter driver writes complete model checkpoints directly, so it has no merge/push options (its --lora_checkpoint_dir argument is reinterpreted internally as the plain checkpoint directory).
Warning
Merging a LoRA adapter into the base weights can slightly degrade measurement accuracy versus serving base + adapter. Keep the unmerged adapter around if you care about the last decimal.
Entry points and other model families#
CoT drivers ship for four families, each with a LoRA and a full-parameter module under medvision_bm.sft:
Family ( |
LoRA driver |
Full-parameter driver |
|---|---|---|
|
|
|
|
|
|
|
|
|
|
|
|
They share the preparation, sampler, and trainer plumbing in medvision_bm.sft.sft_utils (prepare_dataset, prepare_trainer, prepare_trainer_fullFT) and differ mainly in their collate function and chat template. Extending the same recipe to a new family follows the identical two-recipe pattern: a train__SFT-CoT__<family> / train__fullFT-CoT__<family> module reusing these helpers, the matching --model_family_name, the family’s decoder-layer class in --fsdp_transformer_layer_cls_to_wrap, and the right --lmms_eval_opt_deps for env_setup. See Add a model for that walkthrough.
See also#
Reinforcement fine-tuning (RFT) — GRPO-based training on the same tasks.
CLI reference — the full flag list for each entry point.
API reference —
sft_utilsfunctions (prepare_dataset,prepare_trainer,prepare_trainer_fullFT).