unsloth/scripts/benchmarks/qwen3_grpo_unified.py
2026-04-20 13:54:23 +00:00

394 lines
13 KiB
Python

"""Unified entrypoint for Qwen3-4B GRPO backend comparison.
Single script, N backends. Identical dataset / reward functions / sampling /
callbacks so per-step loss, reward, KL, and grad-norm arrays are directly
comparable across runs.
Backends (pick one via `--backend`):
vllm : Unsloth fast_inference=True (vLLM colocated).
unsloth_fi_false : Unsloth fast_inference=False (custom HF inference
kernels + cached fp16 LoRA in fast_linear_forward).
Uses trainer's default (non-vLLM, non-CB) rollout path.
cb_paged : Vanilla HF + PEFT LoRA + transformers continuous
batching with `attn_implementation="paged_attention"`
(FA4 shim active).
cb_sdpa : Same but with `attn_implementation="sdpa_paged"`.
naive_trl : Vanilla HF + PEFT LoRA, no CB, no vLLM (TRL's naive
generate path).
Run:
CUDA_VISIBLE_DEVICES=6 python scripts/benchmarks/qwen3_grpo_unified.py \
--backend vllm --max_steps 10 \
--output_dir outputs/grpo_vllm_10 \
--stats_path logs/grpo_vllm_10.json
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import time
from pathlib import Path
HERE = Path(__file__).resolve().parent
WORKSPACE_ROOT = Path("/mnt/disks/unslothai/ubuntu/workspace_31")
for p in (HERE, WORKSPACE_ROOT):
sys.path.insert(0, str(p))
os.environ.setdefault("UNSLOTH_VLLM_STANDBY", "1")
def parse_args():
p = argparse.ArgumentParser()
p.add_argument(
"--backend",
choices = ["vllm", "unsloth_fi_false", "cb_paged", "cb_sdpa", "naive_trl"],
required = True,
)
p.add_argument("--model_name", default = "unsloth/Qwen3-4B-Base")
p.add_argument("--max_seq_length", type = int, default = 2048)
p.add_argument("--lora_rank", type = int, default = 32)
p.add_argument("--max_steps", type = int, default = 10)
p.add_argument("--num_generations", type = int, default = 4)
p.add_argument("--per_device_train_batch_size", type = int, default = 1)
p.add_argument("--gradient_accumulation_steps", type = int, default = 1)
p.add_argument("--gpu_memory_utilization", type = float, default = 0.75)
p.add_argument("--temperature", type = float, default = 0.1)
p.add_argument("--top_p", type = float, default = 0.97)
p.add_argument("--min_p", type = float, default = 0.5)
p.add_argument("--top_k", type = int, default = 5)
p.add_argument("--learning_rate", type = float, default = 5e-6)
p.add_argument("--max_batch_tokens", type = int, default = 8192)
p.add_argument("--num_blocks", type = int, default = 8192)
p.add_argument("--persistent_cb", action = "store_true")
p.add_argument("--output_dir", required = True)
p.add_argument("--stats_path", required = True)
p.add_argument("--seed", type = int, default = 3407)
return p.parse_args()
def _prepare_common(args):
"""Dataset + rewards are the same for every backend. Always uses the
shared chat template and reward funcs from unsloth_grpo_common."""
from unsloth_grpo_common import (
apply_chat_template_to_tokenizer,
build_dataset,
build_reward_funcs,
build_grpo_kwargs,
)
return (
apply_chat_template_to_tokenizer,
build_dataset,
build_reward_funcs,
build_grpo_kwargs,
)
def _make_stats_callback():
"""StatisticsCallback from torch_debugging_utils. Logs per-step loss,
grad-norm, memory, and wall time. Reward/KL are picked up from the TRL
log dict via `on_log`."""
from torch_debugging_utils import StatisticsCallback
return StatisticsCallback(
track_loss = True,
track_grad_norm = True,
track_memory = True,
track_tensor_stats = False,
)
def _maybe_shim_guided_decoding():
"""Newer vLLM releases have moved GuidedDecodingParams out of
`vllm.sampling_params`; TRL's GRPOTrainer still tries to import it on
the transformers-paged path. Inject a no-op shim if missing."""
try:
import vllm.sampling_params as sp
if not hasattr(sp, "GuidedDecodingParams"):
class _Shim:
def __init__(self, *a, **kw):
pass
sp.GuidedDecodingParams = _Shim
except ImportError:
pass
def _load_unsloth(args, fast_inference: bool):
from unsloth import FastLanguageModel
model, tokenizer = FastLanguageModel.from_pretrained(
model_name = args.model_name,
max_seq_length = args.max_seq_length,
load_in_4bit = False,
fast_inference = fast_inference,
max_lora_rank = args.lora_rank,
**(
{"gpu_memory_utilization": args.gpu_memory_utilization}
if fast_inference
else {}
),
)
model = FastLanguageModel.get_peft_model(
model,
r = args.lora_rank,
target_modules = [
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
],
lora_alpha = args.lora_rank * 2,
use_gradient_checkpointing = "unsloth",
random_state = args.seed,
)
return model, tokenizer
def _load_vanilla_hf(args, attn_impl: str):
"""Vanilla HF + PEFT LoRA. Used by cb_paged / cb_sdpa / naive_trl."""
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig, get_peft_model
tokenizer = AutoTokenizer.from_pretrained(args.model_name)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(
args.model_name,
dtype = torch.bfloat16,
attn_implementation = attn_impl,
).to("cuda")
lora = LoraConfig(
r = args.lora_rank,
lora_alpha = args.lora_rank * 2,
target_modules = [
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
],
bias = "none",
task_type = "CAUSAL_LM",
)
model = get_peft_model(model, lora)
try:
model.gradient_checkpointing_enable(
gradient_checkpointing_kwargs = {"use_reentrant": False}
)
except TypeError:
model.gradient_checkpointing_enable()
model.enable_input_require_grads()
return model, tokenizer
def main():
args = parse_args()
os.makedirs(args.output_dir, exist_ok = True)
os.makedirs(os.path.dirname(os.path.abspath(args.stats_path)) or ".", exist_ok = True)
import torch
from torch_debugging_utils import set_all_seeds_fast
set_all_seeds_fast(args.seed)
# FA4 shim lives here so CB paths dispatch to Blackwell kernels.
import flash_attn_fa4_shim # noqa: F401
flash_attn_fa4_shim.apply()
_maybe_shim_guided_decoding()
(
apply_chat_template_to_tokenizer,
build_dataset,
build_reward_funcs,
build_grpo_kwargs,
) = _prepare_common(args)
# --- load model / tokenizer per backend -----------------------------------
persistent_teardown_target = None
if args.backend == "vllm":
model, tokenizer = _load_unsloth(args, fast_inference = True)
elif args.backend == "unsloth_fi_false":
model, tokenizer = _load_unsloth(args, fast_inference = False)
elif args.backend == "cb_paged":
model, tokenizer = _load_vanilla_hf(args, attn_impl = "paged_attention")
elif args.backend == "cb_sdpa":
model, tokenizer = _load_vanilla_hf(args, attn_impl = "sdpa_paged")
elif args.backend == "naive_trl":
model, tokenizer = _load_vanilla_hf(args, attn_impl = "sdpa")
else:
raise ValueError(args.backend)
apply_chat_template_to_tokenizer(tokenizer)
dataset, maximum_length = build_dataset(
tokenizer, max_seq_length = args.max_seq_length
)
print(f"[{args.backend}] p90 prompt length = {maximum_length}")
reward_funcs = build_reward_funcs(tokenizer)
# --- GRPOConfig: shared core, backend-specific flags ----------------------
shared = build_grpo_kwargs(
tokenizer,
maximum_length,
max_seq_length = args.max_seq_length,
max_steps = args.max_steps,
num_generations = args.num_generations,
per_device_train_batch_size = args.per_device_train_batch_size,
gradient_accumulation_steps = args.gradient_accumulation_steps,
output_dir = args.output_dir,
)
# Overwrite the equivalence-friendly sampling params.
shared["temperature"] = args.temperature
shared["top_p"] = args.top_p
shared["min_p"] = args.min_p
# TRL's TopKLogitsWarper rejects -1; accept an int >=0 only.
shared["top_k"] = args.top_k if args.top_k and args.top_k > 0 else None
shared["learning_rate"] = args.learning_rate
from trl import GRPOConfig, GRPOTrainer
if args.backend == "vllm":
from vllm import SamplingParams
vllm_sp = SamplingParams(
temperature = args.temperature,
top_p = args.top_p,
min_p = args.min_p,
top_k = args.top_k,
seed = args.seed,
stop = [tokenizer.eos_token],
include_stop_str_in_output = True,
)
training_args = GRPOConfig(
use_vllm = True,
vllm_mode = "colocate",
vllm_sampling_params = vllm_sp,
vllm_gpu_memory_utilization = args.gpu_memory_utilization,
**shared,
)
elif args.backend == "unsloth_fi_false":
# Trainer's default rollout path: model.generate. Unsloth's
# fast_inference=False + for_inference() wires the fast single-token
# decode + cached fp16 LoRA.
training_args = GRPOConfig(
use_vllm = False,
bf16 = True,
**shared,
)
elif args.backend in ("cb_paged", "cb_sdpa"):
training_args = GRPOConfig(
use_vllm = False,
use_transformers_paged = True,
bf16 = True,
generation_kwargs = {
"max_batch_tokens": args.max_batch_tokens,
"num_blocks": args.num_blocks,
},
**shared,
)
else: # naive_trl
training_args = GRPOConfig(
use_vllm = False,
bf16 = True,
**shared,
)
stats_cb = _make_stats_callback()
trainer = GRPOTrainer(
model = model,
processing_class = tokenizer,
reward_funcs = reward_funcs,
args = training_args,
train_dataset = dataset,
callbacks = [stats_cb],
)
if args.persistent_cb and args.backend in ("cb_paged", "cb_sdpa"):
from persistent_cb import install_for_model, teardown
base = (
trainer.model_wrapped.base_model.model
if hasattr(trainer.model_wrapped, "base_model")
else trainer.model_wrapped
)
install_for_model(base, trainer.generation_config)
persistent_teardown_target = base
torch.cuda.reset_peak_memory_stats()
t_start = time.perf_counter()
try:
trainer.train()
finally:
if persistent_teardown_target is not None:
from persistent_cb import teardown
teardown(persistent_teardown_target)
train_wall = time.perf_counter() - t_start
stats_cb.save_logs(args.stats_path)
times = [l["time_ms"] for l in stats_cb.logs if "time_ms" in l]
losses = [l["loss"] for l in stats_cb.logs if "loss" in l]
rewards = [l.get("reward") for l in stats_cb.logs if "reward" in l]
kls = [l.get("kl") for l in stats_cb.logs if "kl" in l]
grad_norms = [l.get("grad_norm") for l in stats_cb.logs if "grad_norm" in l]
# Post-warmup (skip first 3 steps) median.
median_step_ms = None
if len(times) > 3:
post = sorted(times[3:])
median_step_ms = post[len(post) // 2]
summary = {
"backend": args.backend,
"max_steps": args.max_steps,
"train_wall_s": train_wall,
"median_step_ms_post_warmup": median_step_ms,
"n_logged_steps": len(stats_cb.logs),
"sampling": {
"temperature": args.temperature,
"top_p": args.top_p,
"min_p": args.min_p,
"top_k": args.top_k,
},
"losses": losses,
"rewards": rewards,
"kls": kls,
"grad_norms": grad_norms,
"step_times_ms": times,
"peak_memory_gb": torch.cuda.max_memory_allocated() / 1024**3,
"logs_path": args.stats_path,
}
summary_path = Path(args.stats_path).with_suffix(".summary.json")
with open(summary_path, "w") as f:
json.dump(summary, f, indent = 2)
print(
json.dumps(
{
k: v
for k, v in summary.items()
if k not in ("losses", "rewards", "kls", "grad_norms", "step_times_ms")
},
indent = 2,
)
)
print(f"\n[{args.backend}] wrote summary to {summary_path}")
# vLLM engine holds refs; fast-exit rather than wait for shutdown.
os._exit(0)
if __name__ == "__main__":
main()