23f2002275 Claude Sonnet 4.6 commited on
Commit ·
d92866e
1
Parent(s): 1bf8189
fix(grpo): allow FATHOM_USE_VLLM=0 to bypass vLLM rollout and avoid IS-ratio collapse under QLoRA
Browse filesGate use_vllm, vllm_mode, and vllm_gpu_memory_utilization on the
FATHOM_USE_VLLM env var (default "1" = enabled). When set to "0",
TRL falls back to HF generate() for rollouts — ~3x slower per step
but importance_sampling_ratio stays near 1.0 instead of 1e-7 to 1e-5.
Root cause of the 200-step flat-reward run: QLoRA + vLLM colocate
merge/unmerge produces token-prob drift between the rollout model
(vLLM, merged 4-bit) and the gradient model (HF, LoRA over 4-bit base),
causing TRL IS clipping to zero every gradient update.
The vllm_mode='colocate' assert is now inside the if use_vllm: branch.
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
- train/grpo.py +23 -12
train/grpo.py
CHANGED
|
@@ -75,10 +75,15 @@ def run_grpo(
|
|
| 75 |
sft_adapter_dir,
|
| 76 |
)
|
| 77 |
|
| 78 |
-
# TRN-03 step 2:
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 82 |
|
| 83 |
# TRN-03 step 2: Build GRPOConfig from cfg.train.
|
| 84 |
# TRL minor versions have changed some GRPOConfig field names; select only
|
|
@@ -115,12 +120,20 @@ def run_grpo(
|
|
| 115 |
_set_if_supported(["max_steps"], int(cfg.train.max_steps))
|
| 116 |
_set_if_supported(["save_steps"], int(cfg.train.save_steps))
|
| 117 |
_set_if_supported(["seed"], int(cfg.seed))
|
| 118 |
-
_set_if_supported(["
|
| 119 |
-
|
| 120 |
-
["
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 124 |
wb_key = os.environ.get("WANDB_API_KEY", "").strip()
|
| 125 |
use_wandb = len(wb_key) >= 40 and wb_key.isalnum()
|
| 126 |
_set_if_supported(["report_to"], ["wandb"] if use_wandb else [])
|
|
@@ -131,8 +144,6 @@ def run_grpo(
|
|
| 131 |
"W&B logging; metrics still printed to stdout.", len(wb_key)
|
| 132 |
)
|
| 133 |
_set_if_supported(["logging_steps"], 1)
|
| 134 |
-
# vllm_mode='colocate' alone is a no-op without `use_vllm=True` in TRL 1.2.
|
| 135 |
-
_set_if_supported(["use_vllm"], True)
|
| 136 |
# Keep TRL from dropping our extra columns (gold_answer etc.) so they
|
| 137 |
# reach the reward function as kwargs.
|
| 138 |
_set_if_supported(["remove_unused_columns"], False)
|
|
|
|
| 75 |
sft_adapter_dir,
|
| 76 |
)
|
| 77 |
|
| 78 |
+
# TRN-03 step 2: FATHOM_USE_VLLM env var gates vLLM rollout.
|
| 79 |
+
# Set FATHOM_USE_VLLM=0 to use HF generate() for rollouts (slower but
|
| 80 |
+
# avoids IS-ratio collapse under QLoRA + vLLM merge drift).
|
| 81 |
+
use_vllm_env = os.environ.get("FATHOM_USE_VLLM", "1").strip()
|
| 82 |
+
use_vllm = use_vllm_env not in ("0", "false", "False", "")
|
| 83 |
+
if use_vllm:
|
| 84 |
+
assert str(cfg.train.vllm_mode) == "colocate", (
|
| 85 |
+
"TRN-03 gate: vllm_mode must be 'colocate' for multi-turn OpenEnv (STACK §10.4)"
|
| 86 |
+
)
|
| 87 |
|
| 88 |
# TRN-03 step 2: Build GRPOConfig from cfg.train.
|
| 89 |
# TRL minor versions have changed some GRPOConfig field names; select only
|
|
|
|
| 120 |
_set_if_supported(["max_steps"], int(cfg.train.max_steps))
|
| 121 |
_set_if_supported(["save_steps"], int(cfg.train.save_steps))
|
| 122 |
_set_if_supported(["seed"], int(cfg.seed))
|
| 123 |
+
_set_if_supported(["use_vllm"], use_vllm)
|
| 124 |
+
if use_vllm:
|
| 125 |
+
_set_if_supported(["vllm_mode"], str(cfg.train.vllm_mode))
|
| 126 |
+
_set_if_supported(
|
| 127 |
+
["vllm_gpu_memory_utilization"],
|
| 128 |
+
float(cfg.train.vllm_gpu_memory_utilization),
|
| 129 |
+
)
|
| 130 |
+
log.info("TRN-03 vLLM rollout enabled (vllm_mode=%s)", cfg.train.vllm_mode)
|
| 131 |
+
else:
|
| 132 |
+
log.warning(
|
| 133 |
+
"TRN-03 FATHOM_USE_VLLM=0 — using HF generate() for "
|
| 134 |
+
"rollouts. Slower but importance_sampling_ratio stays "
|
| 135 |
+
"near 1.0 (avoids QLoRA + vLLM merge drift bug)."
|
| 136 |
+
)
|
| 137 |
wb_key = os.environ.get("WANDB_API_KEY", "").strip()
|
| 138 |
use_wandb = len(wb_key) >= 40 and wb_key.isalnum()
|
| 139 |
_set_if_supported(["report_to"], ["wandb"] if use_wandb else [])
|
|
|
|
| 144 |
"W&B logging; metrics still printed to stdout.", len(wb_key)
|
| 145 |
)
|
| 146 |
_set_if_supported(["logging_steps"], 1)
|
|
|
|
|
|
|
| 147 |
# Keep TRL from dropping our extra columns (gold_answer etc.) so they
|
| 148 |
# reach the reward function as kwargs.
|
| 149 |
_set_if_supported(["remove_unused_columns"], False)
|