From 0557f9151c85242ab8bbbc4c01a201d6ea6e107a Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 20 Apr 2026 14:01:31 +0000 Subject: [PATCH] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- scripts/benchmarks/cb_sync_driver.py | 130 +++++---- scripts/benchmarks/cb_vs_vllm_generation.py | 211 ++++++++------ scripts/benchmarks/make_lora_adapter.py | 46 ++-- scripts/benchmarks/qwen3_grpo_notebook.py | 289 +++++++++++--------- scripts/benchmarks/qwen3_grpo_unified.py | 281 +++++++++++-------- 5 files changed, 568 insertions(+), 389 deletions(-) diff --git a/scripts/benchmarks/cb_sync_driver.py b/scripts/benchmarks/cb_sync_driver.py index f8c2d04492..afafaa4705 100644 --- a/scripts/benchmarks/cb_sync_driver.py +++ b/scripts/benchmarks/cb_sync_driver.py @@ -55,6 +55,7 @@ from transformers.generation.continuous_batching.scheduler import FIFOScheduler @dataclass class CBSyncConfig: """Tunables for the sync driver.""" + max_new_tokens: int = 512 use_cuda_graph: bool = True # Number of eager warmup steps before capturing a CUDA graph. @@ -67,7 +68,7 @@ class CBSyncConfig: max_batch_tokens: int = 8192 num_blocks: int = 8192 # Progress callback (step_index, tokens_produced_total) -> None. - on_step: Optional[callable] = field(default=None) + on_step: Optional[callable] = field(default = None) class SyncCBDriver: @@ -79,8 +80,12 @@ class SyncCBDriver: finished. """ - def __init__(self, model: torch.nn.Module, generation_config: GenerationConfig, - cfg: CBSyncConfig): + def __init__( + self, + model: torch.nn.Module, + generation_config: GenerationConfig, + cfg: CBSyncConfig, + ): self.model = model.eval() self.cfg = cfg # Force-greedy + upper-bound overrides on a copy. @@ -100,11 +105,11 @@ class SyncCBDriver: # We reuse the Manager's methods but never call `.start()`. Its # constructor builds: logit processor, do_sample flag, etc. self.manager = ContinuousBatchingManager( - model=self.model, - generation_config=gc, - manual_eviction=False, - streaming=False, - slice_inputs=False, # fixed-shape views -> CUDA-graph safe + model = self.model, + generation_config = gc, + manual_eviction = False, + streaming = False, + slice_inputs = False, # fixed-shape views -> CUDA-graph safe ) # The manager's `use_cuda_graph` is checked inside `warmup()`, but its # `__init__` refuses to set it. Set it directly now that we bypass @@ -118,7 +123,7 @@ class SyncCBDriver: gc, self.model.device, self.model.dtype, - tp_size=getattr(self.model, "_tp_size", None), + tp_size = getattr(self.model, "_tp_size", None), ) self.batch_processor = ContinuousBatchProcessor( self.cache, @@ -129,10 +134,10 @@ class SyncCBDriver: self.manager.stop_event, self.model.device, self.model.dtype, - FIFOScheduler(self.cache, manual_eviction=False), - streaming=False, - manual_eviction=False, - slice_inputs=False, + FIFOScheduler(self.cache, manual_eviction = False), + streaming = False, + manual_eviction = False, + slice_inputs = False, ) self.manager.batch_processor = self.batch_processor self._graph: Optional[torch.cuda.CUDAGraph] = None @@ -148,13 +153,13 @@ class SyncCBDriver: for _ in range(self.cfg.warmup_steps): self.manager._generation_step(self.batch_processor) torch.cuda.synchronize() - stream = torch.cuda.Stream(device=self.model.device) + stream = torch.cuda.Stream(device = self.model.device) stream.wait_stream(torch.cuda.current_stream()) with torch.cuda.stream(stream): self.manager._generation_step(self.batch_processor) torch.cuda.current_stream().wait_stream(stream) self._graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(self._graph, stream=stream): + with torch.cuda.graph(self._graph, stream = stream): self.manager._generation_step(self.batch_processor) else: self._graph.replay() @@ -206,14 +211,20 @@ class SyncCBDriver: Shape consistency between decodes is what makes the graph replayable. """ try: - return (self.batch_processor.total_query_length - == self.batch_processor.total_batch_size) + return ( + self.batch_processor.total_query_length + == self.batch_processor.total_batch_size + ) except Exception: return False def _produced(self) -> int: - return sum(len(r.generated_tokens) for r - in getattr(self.batch_processor.scheduler, "active_requests", {}).values()) + return sum( + len(r.generated_tokens) + for r in getattr( + self.batch_processor.scheduler, "active_requests", {} + ).values() + ) def close(self): # Caches hold GPU memory; free them explicitly. @@ -223,9 +234,12 @@ class SyncCBDriver: self.manager.batch_processor = None -def cb_sync_generate(model: torch.nn.Module, generation_config: GenerationConfig, - prompt_ids_list: list[list[int]], - cfg: CBSyncConfig) -> dict[str, list[int]]: +def cb_sync_generate( + model: torch.nn.Module, + generation_config: GenerationConfig, + prompt_ids_list: list[list[int]], + cfg: CBSyncConfig, +) -> dict[str, list[int]]: """One-shot entrypoint: build a driver, submit, drain, close. Matches the semantics of `model.generate_batch(...)` but on the main @@ -251,17 +265,18 @@ if __name__ == "__main__": sys.path.insert(0, str(HERE)) import flash_attn_fa4_shim # noqa: E402 + flash_attn_fa4_shim.apply() parser = argparse.ArgumentParser() - parser.add_argument("--model_name", default="unsloth/Qwen3-4B-Base") - parser.add_argument("--n_prompts", type=int, default=32) - parser.add_argument("--max_new_tokens", type=int, default=512) - parser.add_argument("--attn_impl", default="paged_attention") - parser.add_argument("--use_cuda_graph", action="store_true") - parser.add_argument("--max_batch_tokens", type=int, default=8192) - parser.add_argument("--num_blocks", type=int, default=8192) - parser.add_argument("--stats_path", required=True) + parser.add_argument("--model_name", default = "unsloth/Qwen3-4B-Base") + parser.add_argument("--n_prompts", type = int, default = 32) + parser.add_argument("--max_new_tokens", type = int, default = 512) + parser.add_argument("--attn_impl", default = "paged_attention") + parser.add_argument("--use_cuda_graph", action = "store_true") + parser.add_argument("--max_batch_tokens", type = int, default = 8192) + parser.add_argument("--num_blocks", type = int, default = 8192) + parser.add_argument("--stats_path", required = True) args = parser.parse_args() from transformers import AutoModelForCausalLM, AutoTokenizer, GenerationConfig @@ -270,36 +285,49 @@ if __name__ == "__main__": if tok.pad_token is None: tok.pad_token = tok.eos_token model = AutoModelForCausalLM.from_pretrained( - args.model_name, dtype=torch.bfloat16, - attn_implementation=args.attn_impl, + args.model_name, + dtype = torch.bfloat16, + attn_implementation = args.attn_impl, ).to("cuda") model.eval() from unsloth_grpo_common import ( - SYSTEM_PROMPT, apply_chat_template_to_tokenizer, + SYSTEM_PROMPT, + apply_chat_template_to_tokenizer, ) from datasets import load_dataset + apply_chat_template_to_tokenizer(tok) - ds = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split="train") - ds = ds.shuffle(seed=3407).select(range(args.n_prompts)) - messages = [[{"role": "system", "content": SYSTEM_PROMPT}, - {"role": "user", "content": x["prompt"]}] for x in ds] - prompt_ids = [tok.apply_chat_template(m, add_generation_prompt=True, tokenize=True) - for m in messages] + ds = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split = "train") + ds = ds.shuffle(seed = 3407).select(range(args.n_prompts)) + messages = [ + [ + {"role": "system", "content": SYSTEM_PROMPT}, + {"role": "user", "content": x["prompt"]}, + ] + for x in ds + ] + prompt_ids = [ + tok.apply_chat_template(m, add_generation_prompt = True, tokenize = True) + for m in messages + ] gc = GenerationConfig( - max_new_tokens=args.max_new_tokens, do_sample=False, - pad_token_id=tok.pad_token_id, bos_token_id=tok.bos_token_id, - eos_token_id=tok.eos_token_id, use_cache=True, + max_new_tokens = args.max_new_tokens, + do_sample = False, + pad_token_id = tok.pad_token_id, + bos_token_id = tok.bos_token_id, + eos_token_id = tok.eos_token_id, + use_cache = True, ) cfg = CBSyncConfig( - max_new_tokens=args.max_new_tokens, - use_cuda_graph=args.use_cuda_graph, - max_batch_tokens=args.max_batch_tokens, - num_blocks=args.num_blocks, - eos_token_id=tok.eos_token_id, - pad_token_id=tok.pad_token_id or tok.eos_token_id, + max_new_tokens = args.max_new_tokens, + use_cuda_graph = args.use_cuda_graph, + max_batch_tokens = args.max_batch_tokens, + num_blocks = args.num_blocks, + eos_token_id = tok.eos_token_id, + pad_token_id = tok.pad_token_id or tok.eos_token_id, ) torch.cuda.reset_peak_memory_stats() @@ -330,7 +358,7 @@ if __name__ == "__main__": "max_new_tokens": args.max_new_tokens, "peak_memory_gb": torch.cuda.max_memory_allocated() / 1024**3, } - os.makedirs(os.path.dirname(os.path.abspath(args.stats_path)) or ".", exist_ok=True) + os.makedirs(os.path.dirname(os.path.abspath(args.stats_path)) or ".", exist_ok = True) with open(args.stats_path, "w") as f: - json.dump(out, f, indent=2) - print(json.dumps(out, indent=2)) + json.dump(out, f, indent = 2) + print(json.dumps(out, indent = 2)) diff --git a/scripts/benchmarks/cb_vs_vllm_generation.py b/scripts/benchmarks/cb_vs_vllm_generation.py index 795a8340fa..fd4fc715bb 100644 --- a/scripts/benchmarks/cb_vs_vllm_generation.py +++ b/scripts/benchmarks/cb_vs_vllm_generation.py @@ -50,8 +50,8 @@ def build_prompts(tokenizer, n_prompts): from datasets import load_dataset apply_chat_template_to_tokenizer(tokenizer) - ds = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split="train") - ds = ds.shuffle(seed=3407).select(range(n_prompts)) + ds = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split = "train") + ds = ds.shuffle(seed = 3407).select(range(n_prompts)) messages = [ [ {"role": "system", "content": SYSTEM_PROMPT}, @@ -60,11 +60,11 @@ def build_prompts(tokenizer, n_prompts): for x in ds ] prompts_text = [ - tokenizer.apply_chat_template(m, add_generation_prompt=True, tokenize=False) + tokenizer.apply_chat_template(m, add_generation_prompt = True, tokenize = False) for m in messages ] prompt_ids = [ - tokenizer.apply_chat_template(m, add_generation_prompt=True, tokenize=True) + tokenizer.apply_chat_template(m, add_generation_prompt = True, tokenize = True) for m in messages ] return prompts_text, prompt_ids @@ -75,35 +75,37 @@ def run_vllm(args): 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=True, - max_lora_rank=32, - gpu_memory_utilization=args.gpu_memory_utilization, + model_name = args.model_name, + max_seq_length = args.max_seq_length, + load_in_4bit = False, + fast_inference = True, + max_lora_rank = 32, + gpu_memory_utilization = args.gpu_memory_utilization, ) prompts_text, prompt_ids = build_prompts(tokenizer, args.n_prompts) lora_request = None if args.lora_adapter: from vllm.lora.request import LoRARequest + lora_request = LoRARequest("fresh", 1, str(Path(args.lora_adapter).resolve())) from vllm import SamplingParams + sp = SamplingParams( - temperature=args.temperature, - top_p=args.top_p, - min_p=args.min_p, - top_k=args.top_k, - seed=3407, - max_tokens=args.max_new_tokens, - stop=[tokenizer.eos_token], - include_stop_str_in_output=True, + temperature = args.temperature, + top_p = args.top_p, + min_p = args.min_p, + top_k = args.top_k, + seed = 3407, + max_tokens = args.max_new_tokens, + stop = [tokenizer.eos_token], + include_stop_str_in_output = True, ) # Warmup on 16 prompts then discard. warmup_text = prompts_text[:16] - _ = model.fast_generate(warmup_text, sampling_params=sp, lora_request=lora_request) + _ = model.fast_generate(warmup_text, sampling_params = sp, lora_request = lora_request) torch.cuda.synchronize() n_prompt_tokens = sum(len(p) for p in prompt_ids) @@ -114,7 +116,7 @@ def run_vllm(args): torch.cuda.synchronize() t0 = time.perf_counter() outputs = model.fast_generate( - prompts_text, sampling_params=sp, lora_request=lora_request + prompts_text, sampling_params = sp, lora_request = lora_request ) torch.cuda.synchronize() wall_times.append(time.perf_counter() - t0) @@ -122,7 +124,11 @@ def run_vllm(args): last_outputs = outputs med = sorted(wall_times)[len(wall_times) // 2] - sample_texts = [o.outputs[0].text[:200] for o in (last_outputs[:3] or [])] if last_outputs else [] + sample_texts = ( + [o.outputs[0].text[:200] for o in (last_outputs[:3] or [])] + if last_outputs + else [] + ) return { "backend": "vllm", "lora_adapter": args.lora_adapter, @@ -151,16 +157,17 @@ def run_tpaged(args): tokenizer.pad_token = tokenizer.eos_token model = AutoModelForCausalLM.from_pretrained( args.model_name, - dtype=torch.bfloat16, - attn_implementation=args.attn_impl, + dtype = torch.bfloat16, + attn_implementation = args.attn_impl, ).to("cuda") model.eval() if args.lora_adapter: from peft import PeftModel + # NOTE: no merge_adapter -- we measure LoRA-active inference. model = PeftModel.from_pretrained( - model, str(Path(args.lora_adapter).resolve()), is_trainable=False + model, str(Path(args.lora_adapter).resolve()), is_trainable = False ) model.eval() @@ -169,16 +176,16 @@ def run_tpaged(args): prompts_text, prompt_ids = build_prompts(tokenizer, args.n_prompts) gen_config = GenerationConfig( - max_new_tokens=args.max_new_tokens, - do_sample=True, - temperature=args.temperature, - top_p=args.top_p, - min_p=args.min_p, - top_k=args.top_k, - pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id, - bos_token_id=tokenizer.bos_token_id, - eos_token_id=tokenizer.eos_token_id, - use_cache=True, + max_new_tokens = args.max_new_tokens, + do_sample = True, + temperature = args.temperature, + top_p = args.top_p, + min_p = args.min_p, + top_k = args.top_k, + pad_token_id = tokenizer.pad_token_id or tokenizer.eos_token_id, + bos_token_id = tokenizer.bos_token_id, + eos_token_id = tokenizer.eos_token_id, + use_cache = True, ) gen_config.max_batch_tokens = args.max_batch_tokens gen_config.num_blocks = args.num_blocks @@ -188,7 +195,9 @@ def run_tpaged(args): warmup_ids = prompt_ids[:16] with torch.inference_mode(): - _ = model.generate_batch(warmup_ids, generation_config=gen_config, progress_bar=False) + _ = model.generate_batch( + warmup_ids, generation_config = gen_config, progress_bar = False + ) torch.cuda.synchronize() n_prompt_tokens = sum(len(p) for p in prompt_ids) @@ -200,7 +209,7 @@ def run_tpaged(args): t0 = time.perf_counter() with torch.inference_mode(): outputs = model.generate_batch( - prompt_ids, generation_config=gen_config, progress_bar=False + prompt_ids, generation_config = gen_config, progress_bar = False ) torch.cuda.synchronize() wall_times.append(time.perf_counter() - t0) @@ -212,7 +221,7 @@ def run_tpaged(args): if last_outputs is not None: for k in list(last_outputs.keys())[:3]: toks = last_outputs[k].generated_tokens - sample_texts.append(tokenizer.decode(toks, skip_special_tokens=False)[:200]) + sample_texts.append(tokenizer.decode(toks, skip_special_tokens = False)[:200]) med = sorted(wall_times)[len(wall_times) // 2] return { @@ -248,33 +257,40 @@ def run_unsloth_fi_false(args): 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=False, - max_lora_rank=32, + model_name = args.model_name, + max_seq_length = args.max_seq_length, + load_in_4bit = False, + fast_inference = False, + max_lora_rank = 32, ) # Attach LoRA rank 32 the same way the GRPO notebook does. model = FastLanguageModel.get_peft_model( model, - r=32, - target_modules=[ - "q_proj", "k_proj", "v_proj", "o_proj", - "gate_proj", "up_proj", "down_proj", + r = 32, + target_modules = [ + "q_proj", + "k_proj", + "v_proj", + "o_proj", + "gate_proj", + "up_proj", + "down_proj", ], - lora_alpha=64, - use_gradient_checkpointing="unsloth", - random_state=3407, + lora_alpha = 64, + use_gradient_checkpointing = "unsloth", + random_state = 3407, ) # Optional: overlay a shared adapter so weights match other backends. if args.lora_adapter: from safetensors import safe_open + adapter_file = Path(args.lora_adapter).resolve() / "adapter_model.safetensors" loaded_tensors = {} - with safe_open(str(adapter_file), framework="pt") as f: + with safe_open(str(adapter_file), framework = "pt") as f: for key in f.keys(): loaded_tensors[key] = f.get_tensor(key) + # Both PEFT and Unsloth's `get_peft_model` produce parameter names with # `base_model.model.` prefix plus `.lora_{A,B}.default.weight`. Build a # normalized (core-path) -> param map, then match by core path only. @@ -282,10 +298,12 @@ def run_unsloth_fi_false(args): n = name for pref in ("base_model.model.", "model."): if n.startswith(pref): - n = n[len(pref):] + n = n[len(pref) :] n = n.replace(".lora_A.default.", ".lora_A.").replace( - ".lora_B.default.", ".lora_B.") + ".lora_B.default.", ".lora_B." + ) return n + own_by_core = {} for n, p in model.named_parameters(): if "lora_" in n: @@ -299,8 +317,10 @@ def run_unsloth_fi_false(args): own.data.copy_(tensor.to(own.device, own.dtype)) matched += 1 break - print(f"[unsloth_fi_false] LoRA weight sync matched {matched} tensors " - f"(out of {len(loaded_tensors)} adapter entries).") + print( + f"[unsloth_fi_false] LoRA weight sync matched {matched} tensors " + f"(out of {len(loaded_tensors)} adapter entries)." + ) FastLanguageModel.for_inference(model) @@ -308,28 +328,29 @@ def run_unsloth_fi_false(args): # `model.generate` accepts batched input_ids; pad to max length. from transformers import GenerationConfig + if tokenizer.padding_side != "left": tokenizer.padding_side = "left" # decoder needs left padding if tokenizer.pad_token_id is None: tokenizer.pad_token = tokenizer.eos_token gen_config = GenerationConfig( - max_new_tokens=args.max_new_tokens, - do_sample=True, - temperature=args.temperature, - top_p=args.top_p, - min_p=args.min_p, - top_k=args.top_k, - pad_token_id=tokenizer.pad_token_id, - bos_token_id=tokenizer.bos_token_id, - eos_token_id=tokenizer.eos_token_id, - use_cache=True, + max_new_tokens = args.max_new_tokens, + do_sample = True, + temperature = args.temperature, + top_p = args.top_p, + min_p = args.min_p, + top_k = args.top_k, + pad_token_id = tokenizer.pad_token_id, + bos_token_id = tokenizer.bos_token_id, + eos_token_id = tokenizer.eos_token_id, + use_cache = True, ) def _batched_generate(texts): - batch = tokenizer(texts, return_tensors="pt", padding=True).to("cuda") + batch = tokenizer(texts, return_tensors = "pt", padding = True).to("cuda") with torch.inference_mode(): - out = model.generate(**batch, generation_config=gen_config) + out = model.generate(**batch, generation_config = gen_config) prompt_len = batch["input_ids"].shape[1] return out, prompt_len @@ -350,7 +371,9 @@ def run_unsloth_fi_false(args): wall_times.append(time.perf_counter() - t0) # Count generated tokens past prompt_len per sequence (subtract any # trailing pad-only tail by comparing against EOS). - total_decoded = int((out_ids[:, prompt_len:] != tokenizer.pad_token_id).sum().item()) + total_decoded = int( + (out_ids[:, prompt_len:] != tokenizer.pad_token_id).sum().item() + ) last_out_ids = out_ids last_prompt_len = prompt_len @@ -358,8 +381,11 @@ def run_unsloth_fi_false(args): sample_texts = [] if last_out_ids is not None: for i in range(min(3, last_out_ids.shape[0])): - sample_texts.append(tokenizer.decode( - last_out_ids[i, last_prompt_len:], skip_special_tokens=False)[:200]) + sample_texts.append( + tokenizer.decode( + last_out_ids[i, last_prompt_len:], skip_special_tokens = False + )[:200] + ) return { "backend": "unsloth_fi_false", @@ -378,30 +404,35 @@ def run_unsloth_fi_false(args): def parse_args(): p = argparse.ArgumentParser() - p.add_argument("--backend", choices=["vllm", "tpaged", "unsloth_fi_false"], 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("--n_prompts", type=int, default=32) - p.add_argument("--n_rounds", type=int, default=2) - p.add_argument("--max_new_tokens", type=int, default=512) - p.add_argument("--gpu_memory_utilization", type=float, default=0.8) - p.add_argument("--attn_impl", default="sdpa") - p.add_argument("--max_batch_tokens", type=int, default=8192) - p.add_argument("--num_blocks", type=int, default=16384) - p.add_argument("--persistent_cb", action="store_true") - p.add_argument("--lora_adapter", default=None, - help="Path to a PEFT adapter (rank 32) applied in every backend.") - 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("--stats_path", required=True) + p.add_argument( + "--backend", choices = ["vllm", "tpaged", "unsloth_fi_false"], 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("--n_prompts", type = int, default = 32) + p.add_argument("--n_rounds", type = int, default = 2) + p.add_argument("--max_new_tokens", type = int, default = 512) + p.add_argument("--gpu_memory_utilization", type = float, default = 0.8) + p.add_argument("--attn_impl", default = "sdpa") + p.add_argument("--max_batch_tokens", type = int, default = 8192) + p.add_argument("--num_blocks", type = int, default = 16384) + p.add_argument("--persistent_cb", action = "store_true") + p.add_argument( + "--lora_adapter", + default = None, + help = "Path to a PEFT adapter (rank 32) applied in every backend.", + ) + 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("--stats_path", required = True) return p.parse_args() def main(): args = parse_args() - os.makedirs(os.path.dirname(os.path.abspath(args.stats_path)) or ".", exist_ok=True) + os.makedirs(os.path.dirname(os.path.abspath(args.stats_path)) or ".", exist_ok = True) torch.cuda.reset_peak_memory_stats() if args.backend == "vllm": @@ -419,8 +450,8 @@ def main(): "top_k": args.top_k, } with open(args.stats_path, "w") as f: - json.dump(out, f, indent=2) - print(json.dumps(out, indent=2)) + json.dump(out, f, indent = 2) + print(json.dumps(out, indent = 2)) os._exit(0) diff --git a/scripts/benchmarks/make_lora_adapter.py b/scripts/benchmarks/make_lora_adapter.py index 82df17a2da..ef646e94e8 100644 --- a/scripts/benchmarks/make_lora_adapter.py +++ b/scripts/benchmarks/make_lora_adapter.py @@ -20,16 +20,16 @@ from pathlib import Path def parse_args(): p = argparse.ArgumentParser() - p.add_argument("--model_name", default="unsloth/Qwen3-4B-Base") - p.add_argument("--output", default="outputs/lora_rank32_fresh") - p.add_argument("--rank", type=int, default=32) + p.add_argument("--model_name", default = "unsloth/Qwen3-4B-Base") + p.add_argument("--output", default = "outputs/lora_rank32_fresh") + p.add_argument("--rank", type = int, default = 32) return p.parse_args() def main(): args = parse_args() out_dir = Path(args.output).resolve() - out_dir.mkdir(parents=True, exist_ok=True) + out_dir.mkdir(parents = True, exist_ok = True) # Use vanilla HF -- PEFT's save_pretrained yields the canonical # adapter_config.json + adapter_model.safetensors that vLLM's LoRARequest @@ -45,16 +45,23 @@ def main(): # bf16 base; we only need structure + save. Keep on CPU to avoid a GPU load # just for `save_pretrained`. print(f"[make_lora_adapter] Loading {args.model_name} on CPU...") - model = AutoModelForCausalLM.from_pretrained(args.model_name, dtype=torch.bfloat16) + model = AutoModelForCausalLM.from_pretrained(args.model_name, dtype = torch.bfloat16) peft_cfg = LoraConfig( - r=args.rank, - lora_alpha=args.rank * 2, - target_modules=["q_proj", "k_proj", "v_proj", "o_proj", - "gate_proj", "up_proj", "down_proj"], - bias="none", - task_type="CAUSAL_LM", - lora_dropout=0.0, + r = args.rank, + lora_alpha = args.rank * 2, + target_modules = [ + "q_proj", + "k_proj", + "v_proj", + "o_proj", + "gate_proj", + "up_proj", + "down_proj", + ], + bias = "none", + task_type = "CAUSAL_LM", + lora_dropout = 0.0, ) peft_model = get_peft_model(model, peft_cfg) peft_model.print_trainable_parameters() @@ -66,26 +73,31 @@ def main(): with torch.no_grad(): for name, p in peft_model.named_parameters(): if "lora_B" in name: - p.normal_(mean=0.0, std=1e-4) + p.normal_(mean = 0.0, std = 1e-4) n_reinit += 1 - print(f"[make_lora_adapter] Reinitialized {n_reinit} lora_B matrices with tiny gaussian.") + print( + f"[make_lora_adapter] Reinitialized {n_reinit} lora_B matrices with tiny gaussian." + ) peft_model.save_pretrained(str(out_dir)) tok.save_pretrained(str(out_dir)) # Sanity: verify safetensors file present and non-trivial. from safetensors import safe_open + st_path = out_dir / "adapter_model.safetensors" n_zero_tensors = 0 n_tensors = 0 - with safe_open(str(st_path), framework="pt") as f: + with safe_open(str(st_path), framework = "pt") as f: for key in f.keys(): t = f.get_tensor(key) n_tensors += 1 if (t == 0).all().item(): n_zero_tensors += 1 - print(f"[make_lora_adapter] Wrote {n_tensors} tensors to {st_path} " - f"({n_zero_tensors} all-zero).") + print( + f"[make_lora_adapter] Wrote {n_tensors} tensors to {st_path} " + f"({n_zero_tensors} all-zero)." + ) print(f"[make_lora_adapter] Adapter saved to {out_dir}") diff --git a/scripts/benchmarks/qwen3_grpo_notebook.py b/scripts/benchmarks/qwen3_grpo_notebook.py index b231acf9e5..de01cd6344 100644 --- a/scripts/benchmarks/qwen3_grpo_notebook.py +++ b/scripts/benchmarks/qwen3_grpo_notebook.py @@ -33,28 +33,31 @@ for p in (HERE, WORKSPACE_ROOT): def parse_args(): p = argparse.ArgumentParser() - p.add_argument("--stats_path", default="logs/notebook_ref_10.json") - p.add_argument("--output_dir", default="outputs/notebook_ref_10") - p.add_argument("--max_steps", type=int, default=10) - 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("--gpu_memory_utilization", type=float, default=0.85) - p.add_argument("--num_generations", type=int, default=4) - p.add_argument("--per_device_train_batch_size", type=int, default=1) - 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("--skip_sft_pre_finetune", action="store_true", - help="Skip the format-priming SFT stage; go straight to GRPO.") + p.add_argument("--stats_path", default = "logs/notebook_ref_10.json") + p.add_argument("--output_dir", default = "outputs/notebook_ref_10") + p.add_argument("--max_steps", type = int, default = 10) + 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("--gpu_memory_utilization", type = float, default = 0.85) + p.add_argument("--num_generations", type = int, default = 4) + p.add_argument("--per_device_train_batch_size", type = int, default = 1) + 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( + "--skip_sft_pre_finetune", + action = "store_true", + help = "Skip the format-priming SFT stage; go straight to GRPO.", + ) return p.parse_args() def main(): args = parse_args() - os.makedirs(os.path.dirname(os.path.abspath(args.stats_path)) or ".", exist_ok=True) - os.makedirs(args.output_dir, exist_ok=True) + os.makedirs(os.path.dirname(os.path.abspath(args.stats_path)) or ".", exist_ok = True) + os.makedirs(args.output_dir, exist_ok = True) # Import order matters: unsloth must come before transformers/trl. os.environ.setdefault("UNSLOTH_VLLM_STANDBY", "1") @@ -62,23 +65,28 @@ def main(): import torch # noqa: E402 model, tokenizer = FastLanguageModel.from_pretrained( - model_name=args.model_name, - max_seq_length=args.max_seq_length, - load_in_4bit=False, - fast_inference=True, - max_lora_rank=args.lora_rank, - gpu_memory_utilization=args.gpu_memory_utilization, + model_name = args.model_name, + max_seq_length = args.max_seq_length, + load_in_4bit = False, + fast_inference = True, + max_lora_rank = args.lora_rank, + gpu_memory_utilization = args.gpu_memory_utilization, ) 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", + 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=3407, + lora_alpha = args.lora_rank * 2, + use_gradient_checkpointing = "unsloth", + random_state = 3407, ) reasoning_start = "" @@ -119,16 +127,29 @@ def main(): import numpy as np if not args.skip_sft_pre_finetune: - sft_ds = load_dataset("unsloth/OpenMathReasoning-mini", split="cot") - sft_df = sft_ds.to_pandas()[["expected_answer", "problem", "generated_solution"]] - is_number = pd.to_numeric(pd.Series(sft_df["expected_answer"]), errors="coerce").notnull() + sft_ds = load_dataset("unsloth/OpenMathReasoning-mini", split = "cot") + sft_df = sft_ds.to_pandas()[ + ["expected_answer", "problem", "generated_solution"] + ] + is_number = pd.to_numeric( + pd.Series(sft_df["expected_answer"]), errors = "coerce" + ).notnull() sft_df = sft_df.iloc[np.where(is_number)[0]] def format_dataset(x): - thoughts = x["generated_solution"].replace("", "").replace("", "").strip() + thoughts = ( + x["generated_solution"] + .replace("", "") + .replace("", "") + .strip() + ) final_prompt = ( - reasoning_start + thoughts + reasoning_end - + solution_start + x["expected_answer"] + solution_end + reasoning_start + + thoughts + + reasoning_end + + solution_start + + x["expected_answer"] + + solution_end ) return [ {"role": "system", "content": system_prompt}, @@ -136,61 +157,69 @@ def main(): {"role": "assistant", "content": final_prompt}, ] - sft_df["Messages"] = sft_df.apply(format_dataset, axis=1) - sft_df["N"] = sft_df["Messages"].apply(lambda m: len(tokenizer.apply_chat_template(m))) + sft_df["Messages"] = sft_df.apply(format_dataset, axis = 1) + sft_df["N"] = sft_df["Messages"].apply( + lambda m: len(tokenizer.apply_chat_template(m)) + ) sft_df = sft_df.loc[sft_df["N"] <= args.max_seq_length / 2].copy() sft_df["text"] = tokenizer.apply_chat_template( - sft_df["Messages"].values.tolist(), tokenize=False + sft_df["Messages"].values.tolist(), tokenize = False ) sft_dataset = Dataset.from_pandas(sft_df) from trl import SFTTrainer, SFTConfig + sft_trainer = SFTTrainer( - model=model, - tokenizer=tokenizer, - train_dataset=sft_dataset, - args=SFTConfig( - dataset_text_field="text", - per_device_train_batch_size=1, - gradient_accumulation_steps=1, - warmup_steps=5, - num_train_epochs=2, - learning_rate=2e-4, - logging_steps=5, - optim="adamw_8bit", - weight_decay=0.001, - lr_scheduler_type="linear", - seed=3407, - report_to="none", - output_dir=os.path.join(args.output_dir, "sft"), + model = model, + tokenizer = tokenizer, + train_dataset = sft_dataset, + args = SFTConfig( + dataset_text_field = "text", + per_device_train_batch_size = 1, + gradient_accumulation_steps = 1, + warmup_steps = 5, + num_train_epochs = 2, + learning_rate = 2e-4, + logging_steps = 5, + optim = "adamw_8bit", + weight_decay = 0.001, + lr_scheduler_type = "linear", + seed = 3407, + report_to = "none", + output_dir = os.path.join(args.output_dir, "sft"), ), ) sft_trainer.train() del sft_dataset, sft_df, sft_ds, sft_trainer torch.cuda.empty_cache() import gc + gc.collect() # --- GRPO stage ----------------------------------------------------------- - dataset = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split="train") - dataset = dataset.map(lambda x: { - "prompt": [ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": x["prompt"]}, - ], - "answer": x["solution"], - }) + dataset = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split = "train") + dataset = dataset.map( + lambda x: { + "prompt": [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": x["prompt"]}, + ], + "answer": x["solution"], + } + ) - solution_end_regex = r"[\s]{0,}" + "(?:" + re.escape(tokenizer.eos_token) + ")?" + solution_end_regex = ( + r"[\s]{0,}" + "(?:" + re.escape(tokenizer.eos_token) + ")?" + ) match_format = re.compile( rf"{reasoning_end}.*?" rf"{solution_start}(.+?){solution_end_regex}" rf"[\s]{{0,}}$", - flags=re.MULTILINE | re.DOTALL, + flags = re.MULTILINE | re.DOTALL, ) match_numbers = re.compile( solution_start + r".*?[\s]{0,}([-]?[\d\.\,]{1,})", - flags=re.MULTILINE | re.DOTALL, + flags = re.MULTILINE | re.DOTALL, ) def match_format_exactly(completions, **kwargs): @@ -262,10 +291,12 @@ def main(): # Filter long prompts. tokenized = dataset.map( - lambda x: {"tokens": tokenizer.apply_chat_template( - x["prompt"], add_generation_prompt=True, tokenize=True - )}, - batched=False, + lambda x: { + "tokens": tokenizer.apply_chat_template( + x["prompt"], add_generation_prompt = True, tokenize = True + ) + }, + batched = False, ) tokenized = tokenized.map(lambda x: {"L": len(x["tokens"])}) maximum_length = int(np.quantile(tokenized["L"], 0.9)) @@ -277,60 +308,63 @@ def main(): max_completion_length = args.max_seq_length - max_prompt_length from vllm import SamplingParams + vllm_sampling_params = SamplingParams( - temperature=args.temperature, - top_p=args.top_p, - min_p=args.min_p, - top_k=args.top_k, - seed=3407, - stop=[tokenizer.eos_token], - include_stop_str_in_output=True, + temperature = args.temperature, + top_p = args.top_p, + min_p = args.min_p, + top_k = args.top_k, + seed = 3407, + stop = [tokenizer.eos_token], + include_stop_str_in_output = True, ) from trl import GRPOConfig, GRPOTrainer + training_args = GRPOConfig( - vllm_sampling_params=vllm_sampling_params, - temperature=args.temperature, - top_p=args.top_p, - top_k=args.top_k, - learning_rate=5e-6, - weight_decay=0.001, - warmup_ratio=0.1, - lr_scheduler_type="linear", - optim="adamw_8bit", - logging_steps=1, - per_device_train_batch_size=args.per_device_train_batch_size, - gradient_accumulation_steps=1, - num_generations=args.num_generations, - max_prompt_length=max_prompt_length, - max_completion_length=max_completion_length, - max_steps=args.max_steps, - save_steps=args.max_steps + 1, - report_to="none", - output_dir=args.output_dir, - seed=3407, + vllm_sampling_params = vllm_sampling_params, + temperature = args.temperature, + top_p = args.top_p, + top_k = args.top_k, + learning_rate = 5e-6, + weight_decay = 0.001, + warmup_ratio = 0.1, + lr_scheduler_type = "linear", + optim = "adamw_8bit", + logging_steps = 1, + per_device_train_batch_size = args.per_device_train_batch_size, + gradient_accumulation_steps = 1, + num_generations = args.num_generations, + max_prompt_length = max_prompt_length, + max_completion_length = max_completion_length, + max_steps = args.max_steps, + save_steps = args.max_steps + 1, + report_to = "none", + output_dir = args.output_dir, + seed = 3407, ) from torch_debugging_utils import StatisticsCallback + stats_cb = StatisticsCallback( - track_loss=True, - track_grad_norm=True, - track_memory=True, - track_tensor_stats=False, # hooks are noisy + slow on GRPO model + track_loss = True, + track_grad_norm = True, + track_memory = True, + track_tensor_stats = False, # hooks are noisy + slow on GRPO model ) trainer = GRPOTrainer( - model=model, - processing_class=tokenizer, - reward_funcs=[ + model = model, + processing_class = tokenizer, + reward_funcs = [ match_format_exactly, match_format_approximately, check_answer, check_numbers, ], - args=training_args, - train_dataset=dataset, - callbacks=[stats_cb], + args = training_args, + train_dataset = dataset, + callbacks = [stats_cb], ) t0 = time.perf_counter() @@ -361,30 +395,39 @@ def main(): "logs_path": args.stats_path, "peak_memory_gb": torch.cuda.max_memory_allocated() / 1024**3, } - print(json.dumps(summary, indent=2)) + print(json.dumps(summary, indent = 2)) # Canonical quick-inference: produce a few generations for the writeup. rollouts = [] try: from vllm import SamplingParams as SP + sp_sample = SP( - temperature=args.temperature, - top_p=args.top_p, - min_p=args.min_p, - top_k=args.top_k, - max_tokens=256, + temperature = args.temperature, + top_p = args.top_p, + min_p = args.min_p, + top_k = args.top_k, + max_tokens = 256, ) probe_prompts = [ - [{"role": "system", "content": system_prompt}, - {"role": "user", "content": "What is the sqrt of 101?"}], - [{"role": "system", "content": system_prompt}, - {"role": "user", "content": "If 3x+7 = 22, what is x?"}], - [{"role": "system", "content": system_prompt}, - {"role": "user", "content": "What is 17 * 13?"}], + [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": "What is the sqrt of 101?"}, + ], + [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": "If 3x+7 = 22, what is x?"}, + ], + [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": "What is 17 * 13?"}, + ], ] - texts = [tokenizer.apply_chat_template(p, add_generation_prompt=True, tokenize=False) - for p in probe_prompts] - outs = model.fast_generate(texts, sampling_params=sp_sample, lora_request=None) + texts = [ + tokenizer.apply_chat_template(p, add_generation_prompt = True, tokenize = False) + for p in probe_prompts + ] + outs = model.fast_generate(texts, sampling_params = sp_sample, lora_request = None) for t, o in zip(texts, outs): rollouts.append({"prompt": t, "completion": o.outputs[0].text}) except Exception as e: diff --git a/scripts/benchmarks/qwen3_grpo_unified.py b/scripts/benchmarks/qwen3_grpo_unified.py index 0e72c634da..83122388ed 100644 --- a/scripts/benchmarks/qwen3_grpo_unified.py +++ b/scripts/benchmarks/qwen3_grpo_unified.py @@ -42,36 +42,40 @@ 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) + 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) # Phase 4: torch.compile on the training forward. - p.add_argument("--compile_mode", default=None, - choices=[None, "default", "reduce-overhead", - "max-autotune-no-cudagraphs"], - help="If set, torch.compile(model.forward, mode=...) after " - "the trainer is built. vllm backend is excluded; the " - "rollout engine owns its own compile pipeline.") - p.add_argument("--compile_dynamic", action="store_true", default=True) + p.add_argument( + "--compile_mode", + default = None, + choices = [None, "default", "reduce-overhead", "max-autotune-no-cudagraphs"], + help = "If set, torch.compile(model.forward, mode=...) after " + "the trainer is built. vllm backend is excluded; the " + "rollout engine owns its own compile pipeline.", + ) + p.add_argument("--compile_dynamic", action = "store_true", default = True) return p.parse_args() @@ -79,10 +83,18 @@ 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, + 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, ) - return apply_chat_template_to_tokenizer, build_dataset, build_reward_funcs, build_grpo_kwargs def _make_stats_callback(): @@ -90,11 +102,12 @@ def _make_stats_callback(): 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, + track_loss = True, + track_grad_norm = True, + track_memory = True, + track_tensor_stats = False, ) @@ -104,10 +117,13 @@ def _maybe_shim_guided_decoding(): 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 @@ -115,22 +131,34 @@ def _maybe_shim_guided_decoding(): 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_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, + 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 @@ -146,21 +174,28 @@ def _load_vanilla_hf(args, attn_impl: str): tokenizer.pad_token = tokenizer.eos_token model = AutoModelForCausalLM.from_pretrained( args.model_name, - dtype=torch.bfloat16, - attn_implementation=attn_impl, + 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", + 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} + gradient_checkpointing_kwargs = {"use_reentrant": False} ) except TypeError: model.gradient_checkpointing_enable() @@ -170,38 +205,46 @@ def _load_vanilla_hf(args, attn_impl: str): 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) + 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) + ( + 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) + model, tokenizer = _load_unsloth(args, fast_inference = True) elif args.backend == "unsloth_fi_false": - model, tokenizer = _load_unsloth(args, fast_inference=False) + model, tokenizer = _load_unsloth(args, fast_inference = False) elif args.backend == "cb_paged": - model, tokenizer = _load_vanilla_hf(args, attn_impl="paged_attention") + 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") + model, tokenizer = _load_vanilla_hf(args, attn_impl = "sdpa_paged") elif args.backend == "naive_trl": - model, tokenizer = _load_vanilla_hf(args, attn_impl="sdpa") + 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) + 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) @@ -209,12 +252,12 @@ def main(): 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, + 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 @@ -225,18 +268,24 @@ def main(): 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, + 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, + 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": @@ -244,16 +293,16 @@ def main(): # fast_inference=False + for_inference() wires the fast single-token # decode + cached fp16 LoRA. training_args = GRPOConfig( - use_vllm=False, - bf16=True, + 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={ + use_vllm = False, + use_transformers_paged = True, + bf16 = True, + generation_kwargs = { "max_batch_tokens": args.max_batch_tokens, "num_blocks": args.num_blocks, }, @@ -261,55 +310,63 @@ def main(): ) else: # naive_trl training_args = GRPOConfig( - use_vllm=False, - bf16=True, + 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], + 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) + + 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 # Phase 4: torch.compile on the training forward. if args.compile_mode and args.backend != "vllm": from torch_debugging_utils import clear_inductor_cache, CompileDebugger + clear_inductor_cache() - CompileDebugger.enable(graph_breaks=True, recompiles=True) + CompileDebugger.enable(graph_breaks = True, recompiles = True) # Raise Dynamo cache limit so dynamic-shape recompiles don't thrash. import torch._dynamo + torch._dynamo.config.cache_size_limit = 128 try: torch._dynamo.config.allow_unspec_int_on_nn_module = True except AttributeError: pass - print(f"[{args.backend}] Compiling trainer.model.forward " - f"(mode={args.compile_mode}, dynamic={args.compile_dynamic})") + print( + f"[{args.backend}] Compiling trainer.model.forward " + f"(mode={args.compile_mode}, dynamic={args.compile_dynamic})" + ) trainer.model.forward = torch.compile( trainer.model.forward, - mode=args.compile_mode, - dynamic=args.compile_dynamic, + mode = args.compile_mode, + dynamic = args.compile_dynamic, ) # Reference model inside TRL's GRPO loop also runs a forward. ref = getattr(trainer, "ref_model", None) if ref is not None: ref.forward = torch.compile( - ref.forward, mode=args.compile_mode, - dynamic=args.compile_dynamic, + ref.forward, + mode = args.compile_mode, + dynamic = args.compile_dynamic, ) torch.cuda.reset_peak_memory_stats() @@ -319,6 +376,7 @@ def main(): finally: if persistent_teardown_target is not None: from persistent_cb import teardown + teardown(persistent_teardown_target) train_wall = time.perf_counter() - t_start @@ -358,10 +416,17 @@ def main(): } 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)) + 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)