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 files

Gate 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>

Files changed (1) hide show
  1. 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: vllm_mode MUST be 'colocate' (STACK §10.4 / TRL #4543)
79
- assert str(cfg.train.vllm_mode) == "colocate", (
80
- "TRN-03 gate: vllm_mode must be 'colocate' for multi-turn OpenEnv (STACK §10.4)"
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(["vllm_mode"], str(cfg.train.vllm_mode))
119
- _set_if_supported(
120
- ["vllm_gpu_memory_utilization"],
121
- float(cfg.train.vllm_gpu_memory_utilization),
122
- )
123
- import os
 
 
 
 
 
 
 
 
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)