Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

ACL26-Medivstagym: Tool-Augmented Medical VLM with ARMed-style GRPO

Goal: a medical VLM agent that uses external tools (segmentation, grounding, classifier, zoom) during reasoning, and is trained end-to-end with SFT → ARMed-style GRPO to produce clinically grounded answers across MCQ, yes/no, and open-ended VQA.


1. TL;DR

Aspect This work MediX-R1 (arXiv 2602.23363)
Open-ended training signal ARMed adaptive multi-metric reward (BLEU-1 + ROUGE-1 + BERTScore + CosSim, with running-buffer thresholds) LLM-as-judge (vLLM YES/NO) + MedEmbed cosine + modality + format
Tool use during reasoning ✅ 6-server tool router (segmentation / grounding / classifier / zoom) over a <tool_call> protocol; up to 3 turns ❌ no tools
Trainable answer format \boxed{} final answer with optional <think> reasoning + interleaved <tool_call> <thinking>...<answer> XML
Backbone Qwen3-VL-8B-Instruct Qwen2.5-VL-7B
Eval metric ARMed paper metrics (BLEU-1 + ROUGE-1 + BERTScore + CosSim) on 9231 expanded test GPT-5 / LLM-as-judge on 4 benchmarks
Headline result (Step 270 vs vanilla Qwen3-VL on 9231, ARMed metric) CLOSED 0.732, OPEN Avg pending full run published numbers in paper

Design question this repo invites: For tool-augmented medical reasoning where the model emits intermediate <tool_call>s, is ARMed's deterministic 4-metric reward (cheap, reproducible, runs in-process) preferable to MediX-R1's LLM-as-judge reward (semantically richer but needs a judge LLM + extra GPU)? See §7.


2. Architecture

                       ┌─────────────────────────┐
                       │ Qwen3-VL-8B (actor)     │
                       └──────────┬──────────────┘
                                  │ generates text + <tool_call>{...}</tool_call>
                                  ▼
                       ┌─────────────────────────┐
                       │  Tool Router (:5800)    │
                       └──────────┬──────────────┘
                                  │ dispatches by tool name
        ┌────────────┬────────────┼────────────┬────────────┐
        ▼            ▼            ▼            ▼            ▼
   biomedclip   biomedparse    medsam2     medsam3   groundingdino-med
   (:7657)      (:7659)        (:7660)     (:7664)   (:7688)
                                                                  + image_zoom_in
                                                                  + agent4k_resolution (:7961)
  • All tool servers expose /get_observation over HTTP and run in their own conda env so they can pin different CUDA/PyTorch stacks.
  • Tool outputs are appended to the rollout as <observation>...</observation> and re-fed to the actor. Up to max_turns=3 tool calls per question.
  • Trainer: verl-tool (FSDP2, async rollout, GRPO).

3. Training pipeline

3.1 Stage 1 — Supervised Fine-Tuning (SFT)

  • Base: Qwen3-VL-8B-Instruct.
  • Data: medical VQA mixture (MCQ + yes/no + open-ended), cleaned and reformatted into Qwen3-VL chat template.
  • Code: VlmGym/SFT/
  • Output: VlmGym/SFT/sft_model/qwen3vl_mcq3_open_en_full_sft/ (used as GRPO init).

3.2 Stage 2 — GRPO with tools (this is the contribution)


4. Reward design (the interesting part)

Per-sample reward depends on the task type detected from the dataset's answer_type field.

4.1 MCQ / CLOSED (rule-based, ±1)

reward = +1 if Format ∧ AnswerMatch else -1
  • Format checks the model emitted exactly one \boxed{...} final answer.
  • AnswerMatch uses letter parsing for MCQ; loose yes/no normalization for CLOSED.

4.2 OPEN (ARMed adaptive semantic reward)

Direct port of the ARMed paper (arXiv 2508.12957), adapted to our \boxed{} format and run on CPU inside Ray actors.

R_c  = λ₁ · BLEU-1 + (1 - λ₁) · ROUGE-1                      # textual correctness
R_as = λ₂ · BERTScore_adapt + (1 - λ₂) · CosSim_adapt        # adaptive semantic
R_f  = format_correct ? 1 : 0                                # format
R_total = (γ₁ · R_c + γ₂ · R_as + γ₃ · R_f) / Σγ             # ∈ [0, 1]
reward  = 2 · R_total - 1                                    # → [-1, +1] to match MCQ/CLOSED

Defaults: λ₁=0.5, λ₂=0.2, γ₁=0.4, γ₂=0.4, γ₃=0.2.

Adaptive thresholds (the ARMed twist). BERTScore and CosSim both have a non-trivial baseline floor (~0.6–0.8) on medical text, so a naive sum saturates fast. ARMed maintains a per-instance running buffer (size 256) of recent BERTScore/CosSim values and centers each score on the running median:

T_t = clip(T_{t-1} ± δ_max, T_min=0.0, T_max=0.95)        # bounded drift toward percentile of buffer
R_adapt(x, T) =
    if x ≥ T : 0.5 + 0.5 · (1 - exp(-α_pos · (x - T)))    # squashed >0
    else     : 0.5 · exp(α_neg · (x - T))                  # squashed <0.5

This keeps the OPEN reward signal alive throughout training instead of plateauing at the embedding-similarity floor.

Backbones used for the reward (lazy-loaded, CPU by default to keep GPU free for rollouts):

  • BERTScore: microsoft/BiomedNLP-PubMedBERT-base-uncased-abstract (layer 8)
  • CosSim: pritamdeka/BioBERT-mnli-snli-scinli-scitail-mednli-stsb

5. Tool ecosystem

Tool Purpose Port Used for
biomedclip Zero-shot medical image classification 7657 "what modality? what region?"
biomedparse Generic medical image parsing / labelmap 7659 structure overview
medsam2 Promptable medical segmentation 7660 "what is in this bbox?"
medsam3 Newer MedSAM variant 7664 same, alt backbone
groundingdino-med Open-vocabulary detection 7688 text-prompted localization
image_zoom_in bbox crop + resize (in-process) "look closer at..."
agent4k_resolution High-res re-render 7961 small-feature questions

All servers expose /get_observation and /health. Start them via verl-tool/start_all_tools.sh and the router via verl-tool/tool_router.py.

Smoke test: tool_smoke_test.sh and test_all_tools.sh.


6. Evaluation

6.1 Datasets

Source OPEN CLOSED MCQ total
SLAKE (EN test) 645 416 – 1061
VQA-RAD 200 251 – 451
VQA-Med-2021 1000 – – 1000
Path-VQA (test) 3357 3362 – 6719
Total 5202 4029 – 9231

Combined manifest: dataset/sub_testdata/expanded_test_data.json.

6.2 Metric — primary: ARMed Avg

ARMed Avg = mean(BLEU-1, ROUGE-1, BERTScore, CosSim)   # OPEN
            accuracy                                    # CLOSED / MCQ

Implementation: VlmGym/SFT/armed_evaluator.py.

6.3 Metric — secondary: GPT-5 lenient LLM-as-judge

Used as a sanity check (not for paper main numbers, since it's not reproducible without API access).

  • Implementation: VlmGym/SFT/gpt5_judge_evaluator.py
  • Prompts GPT-5 with (question, ground_truth, predicted) → JSON {"correct": bool, "reason": str}.
  • Explicitly lenient: accepts cross-language equivalents (中/英), IS-A specialization, partial mention of multi-item GTs.
  • Verified: ARMed and GPT-5-judge agree on CLOSED (rule scoring fine for English-only vanilla outputs), but GPT-5 judge is lower on OPEN than ARMed Avg because BERTScore + CosSim have a non-zero floor that inflates ARMed Avg.

6.4 Baselines compared

Model Tools? Size Status
Qwen3-VL-8B-Instruct (vanilla) ❌ 8B ✅ evaluated on 9231
Qwen2.5-VL-7B-Instruct (vanilla) ❌ 7B ✅ evaluated on 9231
InternVL3-8B (vanilla) ❌ 8B ✅ evaluated on 9231
LLaVA-Med-v1.5-Mistral-7B (vanilla) ❌ 7B ✅ evaluated on 9231
Ours: Qwen3-VL + SFT + ARMed-GRPO + tools ✅ 8B 🟡 Step 270 inference in progress on full 9231

Inference scripts (one per baseline):

6.5 Checkpoint selection (800-sample subtest, ARMed OPEN Avg + CLOSED accuracy)

Checkpoint CLOSED OPEN Avg (ARMed) Selected?
step150 0.585 – no
step180 0.611 0.353 no
step240 0.598 – no
step270 0.611 best on OPEN (subtest) ✅

7. Discussion: ours vs MediX-R1

7.1 Where the two papers agree

  • Both reject pure exact-match / BLEU as the only training signal for OPEN medical VQA.
  • Both keep MCQ / CLOSED rule-based.
  • Both use a format reward.

7.2 Where they diverge

Question This work (ARMed-style) MediX-R1
What rewards open-ended correctness? 4 deterministic metrics ensembled with adaptive thresholds (BLEU-1 + ROUGE-1 + BERTScore + CosSim) LLM judge YES/NO (weight 0.575) + MedEmbed cosine (0.375)
Modality bonus? not used as a separate term yes (weight 0.05)
Tool use as part of the policy? yes — <tool_call> is part of the action space, observations are masked from loss but conditioning no
Reward compute cost per rollout ~CPU only (PubMedBERT + BioBERT + ROUGE); no extra GPU needs an extra LLM judge (vLLM) running alongside
Reproducibility outside Anthropic/OpenAI full depends on judge model availability
Saturation behavior adaptive threshold prevents BERTScore floor from saturating depends on judge calibration

7.3 Which design is better — for which setting?

  • If the workload is agentic (model must call segmentation/grounding/classifier mid-reasoning, like ours): adding a vLLM judge alongside the trainer + the tool servers may not be tractable on a single 8×GPU box. ARMed-style determinism keeps GPUs free for rollouts.
  • If the workload is non-agentic open-ended QA and a strong judge is available cheaply: MediX-R1's LLM-judge ceiling on semantic correctness is probably higher, especially for cases where surface metrics under-credit synonyms (e.g. "thoracic" ≡ "chest").
  • An incremental ablation worth running on this codebase: keep ARMed's deterministic backbone, but add MediX-R1's modality-classification term (zero extra GPU). Hypothesis: small but cheap gain.

A formal head-to-head comparison would require running MediX-R1's exact reward in our medical_reasoning-armed.py slot and re-training from the same SFT initializer — not done yet.


8. Layout

ACL26-Medivstagym/
├── README.md                        ← this file
├── VlmGym/SFT/                      ← SFT scripts + armed_evaluator + gpt5_judge_evaluator
├── verl-tool/                       ← fork of verl with tool-call rollout + reward managers
│   ├── verl_tool/workers/reward_manager/medical_reasoning-armed.py   ← REWARD
│   ├── examples/train/train_qwen3vl_8b_mcq3_open_armed_grpo.sh        ← TRAIN
│   ├── examples/train/medical_reasoning_qwen3vl_8b_mcq3_open_armed_grpo_v2/  ← CKPTS
│   ├── start_all_tools.sh           ← tool router bring-up
│   └── tool_router.py
├── Inference_vanilla_*.py           ← 4 vanilla baseline inference scripts
├── Inference_qwen3_trained_model.py ← tool-augmented inference (Step 270)
├── dataset/sub_testdata/            ← test set manifests (800-subset + 9231 expanded)
└── inference_results/               ← jsonl predictions + per_sample + summary.json
    ├── step270_full_9231.jsonl
    ├── {model}_vanilla_9231.jsonl
    ├── {model}_vanilla_armed_eval/summary.json
    └── {model}_vanilla_gpt5_judge/summary.json

9. Reproducing key numbers

# 1) Bring up tool router + 6 tool servers
cd /data/ACL26-Medivstagym/verl-tool && bash start_all_tools.sh
curl http://localhost:5800/health   # expect 200

# 2) GRPO (assumes SFT model already exists)
bash examples/train/train_qwen3vl_8b_mcq3_open_armed_grpo.sh

# 3) Evaluate a checkpoint on the full 9231 test set
python /data/ACL26-Medivstagym/Inference_qwen3_trained_model.py \
  --model-path ../medical_reasoning_qwen3vl_8b_mcq3_open_armed_grpo_v2/step270_merged \
  --data-path  /data/ACL26-Medivstagym/dataset/sub_testdata/expanded_test_data.json \
  --output     /data/ACL26-Medivstagym/inference_results/step270_full_9231.jsonl \
  --save-trace --resume

# 4) Score with ARMed metric
python VlmGym/SFT/armed_evaluator.py \
  --pred_file inference_results/step270_full_9231.jsonl \
  --output_dir inference_results/step270_full_armed_eval/ \
  --device cpu

# 5) (optional) sanity-check with GPT-5 lenient judge
OPENAI_API_KEY=... python VlmGym/SFT/gpt5_judge_evaluator.py \
  --pred-file inference_results/step270_full_9231.jsonl \
  --test-data dataset/sub_testdata/expanded_test_data.json \
  --output-dir inference_results/step270_full_gpt5_judge/ \
  --workers 25

10. Open questions for collaborators

  1. Should step270_merged be the final checkpoint, or train further (e.g. step330+) once the tool servers are stable enough for long runs?
  2. Path-VQA OPEN is the weakest slice across all baselines (and ours) — is this a model issue, a GT-quality issue (many one-word "respiratory"-style GTs), or a metric issue? Worth deeper error analysis before chasing reward changes.
  3. Is the MediX-R1 modality-reward term worth porting in as a 5% additive bonus? Cheap experiment.
  4. Should the eval table in the paper report both ARMed Avg and GPT-5 lenient judge side-by-side, or only the reproducible ARMed numbers?

About

No description, website, or topics provided.

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages