vadimbelsky commited on
Commit
b5731e9
·
verified ·
1 Parent(s): 332e298

Update model card with full 3-stage training pipeline explanation

Browse files
Files changed (1) hide show
  1. README.md +146 -11
README.md CHANGED
@@ -1,21 +1,156 @@
1
  ---
 
 
2
  tags:
3
- - gguf
4
- - llama.cpp
 
 
 
 
5
  - unsloth
 
6
  - vision-language-model
 
 
7
  ---
8
 
9
- # qwen3.5-medical-ft-stage3-dpo : GGUF
10
 
11
- This model was finetuned and converted to GGUF format using [Unsloth](https://github.com/unslothai/unsloth).
12
 
13
- **Example usage**:
14
- - For text only LLMs: `llama-cli -hf vadimbelsky/qwen3.5-medical-ft-stage3-dpo --jinja`
15
- - For multimodal models: `llama-mtmd-cli -hf vadimbelsky/qwen3.5-medical-ft-stage3-dpo --jinja`
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
16
 
17
- ## Available Model files:
18
- - `merged_stage1.F16.gguf`
19
- - `merged_stage1.BF16-mmproj.gguf`
20
- This was trained 2x faster with [Unsloth](https://github.com/unslothai/unsloth)
21
  [<img src="https://raw.githubusercontent.com/unslothai/unsloth/main/images/unsloth%20made%20with%20love.png" width="200"/>](https://github.com/unslothai/unsloth)
 
 
 
 
 
 
 
1
  ---
2
+ language:
3
+ - en
4
  tags:
5
+ - medical
6
+ - triage
7
+ - emergency-medicine
8
+ - esi
9
+ - dpo
10
+ - rlhf
11
  - unsloth
12
+ - qwen3
13
  - vision-language-model
14
+ base_model: Qwen/Qwen3.5-9B
15
+ license: apache-2.0
16
  ---
17
 
18
+ # qwen3.5-medical-ft-stage3-dpo
19
 
20
+ **Qwen3.5-9B** fine-tuned through a three-stage supervised + preference-alignment pipeline for **emergency department triage** using the [Emergency Severity Index (ESI)](https://www.acep.org/patient-care/esi/) 1–5 scale.
21
 
22
+ This repository contains the **Stage 3** merged 16-bit weights — the final production model.
23
+
24
+ ---
25
+
26
+ ## Training Pipeline
27
+
28
+ ### Stage 1 — Domain SFT (medical Q&A)
29
+ - **Base model**: `Qwen/Qwen3.5-9B`
30
+ - **Dataset**: [`vadimbelsky/medical-triage-qa-50k`](https://huggingface.co/datasets/vadimbelsky/medical-triage-qa-50k) — 50 k medical triage Q&A pairs
31
+ - **Method**: Supervised fine-tuning with LoRA (r=16) via Unsloth + TRL SFTTrainer
32
+ - **Goal**: Inject emergency medicine domain knowledge and ESI reasoning into the base model
33
+
34
+ ### Stage 2 — Continued SFT (intake notes)
35
+ - **Base model**: Stage 1 LoRA merged into base weights
36
+ - **Dataset**: `intake_notes_10k.jsonl` — 10 k real-format SOAP intake notes with structured triage decisions
37
+ - **Method**: Continued SFT with LoRA (r=32, alpha=64) — doubled rank to capture finer-grained triage reasoning
38
+ - **Goal**: Align the model to the exact input/output format used in clinical practice (SOAP note → ESI level + justification + interventions)
39
+
40
+ ### Stage 3 — DPO Alignment (this model)
41
+ - **Base model**: Stage 2 LoRA checkpoint
42
+ - **Dataset**: `dpo_dataset_clean.jsonl` — preference pairs targeting over-triage correction
43
+ - **Method**: Direct Preference Optimization ([DPO](https://arxiv.org/abs/2305.18290)) via TRL DPOTrainer
44
+ - **Goal**: Reduce systematic over-escalation of low-acuity patients (ESI 3/4/5 → ESI 1/2) while preserving 100% recall on genuinely high-risk patients
45
+
46
+ #### DPO hyperparameters
47
+ | Parameter | Value | Notes |
48
+ |-----------|-------|-------|
49
+ | beta (KL penalty) | 0.3 | beta=0.1 caused reward margin to explode to 31, leading to under-triage |
50
+ | Loss type | sigmoid | Standard DPO |
51
+ | Learning rate | 5e-5 | Lower than SFT to avoid catastrophic forgetting |
52
+ | Epochs | 0.15 | DPO overfits fast — loss hits zero well before epoch 1 |
53
+ | Effective batch size | 16 | 2 x device batch x 8 gradient accumulation steps |
54
+ | LoRA r / alpha | 8 / 8 | Conservative rank — DPO requires minimal capacity |
55
+ | Optimizer | AdamW 8-bit | |
56
+ | Precision | BF16 | |
57
+
58
+ ---
59
+
60
+ ## Model Task
61
+
62
+ Given a **SOAP intake note**, the model produces a structured triage decision:
63
+
64
+ - **ESI level** (1–5) with clinical justification
65
+ - **Key clinical findings** driving the decision
66
+ - **Time-to-provider target**
67
+ - **Immediate interventions** required
68
+
69
+ **System prompt**:
70
+ ```
71
+ You are an expert emergency medicine triage nurse. Given a SOAP intake note, provide a structured triage decision including ESI level with justification, key clinical findings, time-to-provider target, and any immediate interventions required.
72
+ ```
73
+
74
+ ### ESI Scale Reference
75
+ | Level | Acuity | Time-to-Provider |
76
+ |-------|--------|-----------------|
77
+ | ESI 1 | Immediate life threat | Immediate |
78
+ | ESI 2 | High risk / emergent | < 10 min |
79
+ | ESI 3 | Urgent, stable | 30–60 min |
80
+ | ESI 4 | Less urgent | 1–2 hours |
81
+ | ESI 5 | Non-urgent | 2–4 hours |
82
+
83
+ ---
84
+
85
+ ## Stage 3 Training Targets
86
+
87
+ | Metric | Target | Pre-DPO Baseline |
88
+ |--------|--------|-----------------|
89
+ | Overall accuracy | > 82% | — |
90
+ | Over-triage rate (ESI 3/4/5 escalated to 1/2) | < 10% | 22.2% |
91
+ | Under-triage rate | < 6% | — |
92
+ | High-risk recall (ESI 1/2) | 100% | 100% |
93
+ | ESI 3 accuracy | > 65% | ~40% |
94
+
95
+ ---
96
+
97
+ ## Files
98
+
99
+ | File | Description |
100
+ |------|-------------|
101
+ | `model.safetensors-0000{1-4}-of-00004.safetensors` | Merged 16-bit weights (4-shard) |
102
+ | `config.json` | Model configuration |
103
+ | `tokenizer.json` / `tokenizer_config.json` | Tokenizer |
104
+ | `processor_config.json` | Vision processor config |
105
+ | `chat_template.jinja` | Chat template |
106
+
107
+ ---
108
+
109
+ ## Usage
110
+
111
+ ```python
112
+ from transformers import AutoModelForCausalLM, AutoTokenizer
113
+
114
+ model = AutoModelForCausalLM.from_pretrained(
115
+ "vadimbelsky/qwen3.5-medical-ft-stage3-dpo",
116
+ torch_dtype="auto",
117
+ device_map="auto",
118
+ )
119
+ tokenizer = AutoTokenizer.from_pretrained("vadimbelsky/qwen3.5-medical-ft-stage3-dpo")
120
+
121
+ SYSTEM_PROMPT = (
122
+ "You are an expert emergency medicine triage nurse. "
123
+ "Given a SOAP intake note, provide a structured triage decision including "
124
+ "ESI level with justification, key clinical findings, time-to-provider target, "
125
+ "and any immediate interventions required."
126
+ )
127
+
128
+ soap_note = """S (Subjective): 58-year-old male with sudden onset chest pain radiating to left arm, diaphoresis, onset 30 minutes ago...
129
+ O (Objective): BP 160/95, HR 110, RR 22, SpO2 94% on room air, Temp 37.1C..."""
130
+
131
+ messages = [
132
+ {"role": "system", "content": SYSTEM_PROMPT},
133
+ {"role": "user", "content": soap_note},
134
+ ]
135
+
136
+ inputs = tokenizer.apply_chat_template(
137
+ messages, tokenize=True, add_generation_prompt=True, return_tensors="pt"
138
+ ).to(model.device)
139
+
140
+ outputs = model.generate(inputs, max_new_tokens=512)
141
+ print(tokenizer.decode(outputs[0][inputs.shape[-1]:], skip_special_tokens=True))
142
+ ```
143
+
144
+ ---
145
+
146
+ ## Training Infrastructure
147
+
148
+ Trained with [Unsloth](https://github.com/unslothai/unsloth) for 2x faster throughput and reduced VRAM usage.
149
 
 
 
 
 
150
  [<img src="https://raw.githubusercontent.com/unslothai/unsloth/main/images/unsloth%20made%20with%20love.png" width="200"/>](https://github.com/unslothai/unsloth)
151
+
152
+ ---
153
+
154
+ ## Disclaimer
155
+
156
+ This model is intended for **research and educational purposes only**. It is not validated for clinical use and must not be used to make real patient triage decisions without oversight from licensed medical professionals.