diff --git a/scripts/benchmarks/cb_vs_vllm_generation.py b/scripts/benchmarks/cb_vs_vllm_generation.py index 8a4f5b6cc7..dcf183e4b6 100644 --- a/scripts/benchmarks/cb_vs_vllm_generation.py +++ b/scripts/benchmarks/cb_vs_vllm_generation.py @@ -32,6 +32,7 @@ import torch # noqa: E402 # No-op for the vLLM backend since vLLM doesn't go through transformers' # attention interface. import flash_attn_fa4_shim # noqa: E402 + flash_attn_fa4_shim.apply() @@ -43,8 +44,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}, @@ -53,11 +54,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 @@ -65,30 +66,36 @@ def build_prompts(tokenizer, n_prompts): def run_vllm(args): import os as _os + _os.environ.setdefault("UNSLOTH_VLLM_STANDBY", "1") 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) from vllm import SamplingParams + sp = SamplingParams( - temperature=1.0, min_p=0.1, top_p=1.0, top_k=-1, - seed=3407, - max_tokens=args.max_new_tokens, - stop=[tokenizer.eos_token], - include_stop_str_in_output=True, + temperature = 1.0, + min_p = 0.1, + top_p = 1.0, + top_k = -1, + 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=None) + _ = model.fast_generate(warmup_text, sampling_params = sp, lora_request = None) torch.cuda.synchronize() # Three measured rounds on the full batch. @@ -98,7 +105,9 @@ def run_vllm(args): for _ in range(args.n_rounds): torch.cuda.synchronize() t0 = time.perf_counter() - outputs = model.fast_generate(prompts_text, sampling_params=sp, lora_request=None) + outputs = model.fast_generate( + prompts_text, sampling_params = sp, lora_request = None + ) torch.cuda.synchronize() wall_times.append(time.perf_counter() - t0) decoded = sum(len(o.outputs[0].token_ids) for o in outputs) @@ -128,8 +137,8 @@ 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() @@ -138,15 +147,15 @@ 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=1.0, - top_p=1.0, - min_p=0.1, - 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 = 1.0, + top_p = 1.0, + min_p = 0.1, + 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, ) # Raise the paged-cache upper bounds; defaults (256 / 4096) throttle CB. gen_config.max_batch_tokens = args.max_batch_tokens @@ -158,7 +167,9 @@ def run_tpaged(args): # Warmup on 16 prompts. 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) @@ -169,7 +180,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) @@ -194,25 +205,28 @@ def run_tpaged(args): def parse_args(): p = argparse.ArgumentParser() - p.add_argument("--backend", choices=["vllm", "tpaged"], 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=64) - p.add_argument("--n_rounds", type=int, default=3) - p.add_argument("--max_new_tokens", type=int, default=1024) - 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", - help="Reuse a single ContinuousBatchingManager across warmup + measured rounds.") - p.add_argument("--stats_path", required=True) + p.add_argument("--backend", choices = ["vllm", "tpaged"], 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 = 64) + p.add_argument("--n_rounds", type = int, default = 3) + p.add_argument("--max_new_tokens", type = int, default = 1024) + 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", + help = "Reuse a single ContinuousBatchingManager across warmup + measured rounds.", + ) + 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": @@ -222,8 +236,8 @@ def main(): out["peak_memory_gb"] = torch.cuda.max_memory_allocated() / 1024**3 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)) # When --persistent_cb is set the background CB worker thread keeps the # process alive. Exit fast; the stats file is already flushed. os._exit(0) diff --git a/scripts/benchmarks/persistent_cb.py b/scripts/benchmarks/persistent_cb.py index 9787ba50c5..8ac9e34b4d 100644 --- a/scripts/benchmarks/persistent_cb.py +++ b/scripts/benchmarks/persistent_cb.py @@ -32,7 +32,9 @@ _ATTR = "_persistent_cb_manager" _LOCK_ATTR = "_persistent_cb_lock" -def install_for_model(model: torch.nn.Module, generation_config: GenerationConfig) -> None: +def install_for_model( + model: torch.nn.Module, generation_config: GenerationConfig +) -> None: """Replace `model.generate_batch` with a version that reuses one manager. The replacement accepts the same arguments as the stock method. A @@ -57,7 +59,11 @@ def install_for_model(model: torch.nn.Module, generation_config: GenerationConfi if not inputs: return {} - gen_config = generation_config or getattr(self, "_persistent_cb_gen_config", None) or self.generation_config + gen_config = ( + generation_config + or getattr(self, "_persistent_cb_gen_config", None) + or self.generation_config + ) lock = getattr(self, _LOCK_ATTR) with lock: @@ -70,15 +76,15 @@ def install_for_model(model: torch.nn.Module, generation_config: GenerationConfi ) if stale: try: - manager.stop(block=True, timeout=5.0) + manager.stop(block = True, timeout = 5.0) except Exception: pass setattr(self, _ATTR, None) manager = None if manager is None: manager = self.init_continuous_batching( - generation_config=gen_config, - slice_inputs=slice_inputs, + generation_config = gen_config, + slice_inputs = slice_inputs, ) manager.start() setattr(self, _ATTR, manager) @@ -88,7 +94,7 @@ def install_for_model(model: torch.nn.Module, generation_config: GenerationConfi manager.add_requests(inputs, **kwargs) finished = 0 while finished < num_requests: - result = manager.get_result(timeout=1) + result = manager.get_result(timeout = 1) if result is None: if not manager.is_running(): break @@ -108,7 +114,7 @@ def teardown(model: torch.nn.Module) -> None: manager = getattr(model, _ATTR, None) if manager is not None: try: - manager.stop(block=True, timeout=5.0) + manager.stop(block = True, timeout = 5.0) except Exception: pass if hasattr(model, "_persistent_cb_original_generate_batch"): diff --git a/scripts/benchmarks/qwen3_grpo_naive.py b/scripts/benchmarks/qwen3_grpo_naive.py index 2a03083314..b6b079b1f1 100644 --- a/scripts/benchmarks/qwen3_grpo_naive.py +++ b/scripts/benchmarks/qwen3_grpo_naive.py @@ -32,10 +32,13 @@ sys.path.insert(0, str(HERE)) # even when vLLM is installed but the GuidedDecodingParams symbol has moved. try: import vllm.sampling_params as _vllm_sp + if not hasattr(_vllm_sp, "GuidedDecodingParams"): + class _GuidedDecodingParamsShim: # pragma: no cover def __init__(self, *a, **kw): pass + _vllm_sp.GuidedDecodingParams = _GuidedDecodingParamsShim except ImportError: pass @@ -54,29 +57,33 @@ from unsloth_grpo_common import ( # noqa: E402 def parse_args(): p = argparse.ArgumentParser() - 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=20) - p.add_argument("--num_generations", type=int, default=2) - p.add_argument("--per_device_train_batch_size", type=int, default=2) - p.add_argument("--gradient_accumulation_steps", type=int, default=1) - p.add_argument("--attn_impl", default="sdpa", - help="Attention implementation: sdpa or flash_attention_2 (FA4 shim installed).") - p.add_argument("--output_dir", default="outputs/grpo_naive") - p.add_argument("--stats_path", default="logs/naive_stats.json") + 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 = 20) + p.add_argument("--num_generations", type = int, default = 2) + p.add_argument("--per_device_train_batch_size", type = int, default = 2) + p.add_argument("--gradient_accumulation_steps", type = int, default = 1) + p.add_argument( + "--attn_impl", + default = "sdpa", + help = "Attention implementation: sdpa or flash_attention_2 (FA4 shim installed).", + ) + p.add_argument("--output_dir", default = "outputs/grpo_naive") + p.add_argument("--stats_path", default = "logs/naive_stats.json") return p.parse_args() 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) # Install the FA4 shim only if the caller asked for flash_attention_2. # For sdpa we leave transformers untouched. if args.attn_impl == "flash_attention_2": import flash_attn_fa4_shim + flash_attn_fa4_shim.apply() tokenizer = AutoTokenizer.from_pretrained(args.model_name) @@ -84,51 +91,61 @@ def main(): 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") 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", + 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", + bias = "none", + task_type = "CAUSAL_LM", ) model = get_peft_model(model, lora) try: - model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}) + model.gradient_checkpointing_enable( + gradient_checkpointing_kwargs = {"use_reentrant": False} + ) except TypeError: model.gradient_checkpointing_enable() model.enable_input_require_grads() 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"[naive] Max prompt length (p90): {maximum_length}") reward_funcs = build_reward_funcs(tokenizer) from trl import GRPOConfig, GRPOTrainer + 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, ) # transformers' TopKLogitsWarper rejects -1. None skips the warper. shared["top_k"] = None training_args = GRPOConfig( - use_vllm=False, - use_transformers_paged=False, - bf16=True, + use_vllm = False, + use_transformers_paged = False, + bf16 = True, **shared, ) @@ -144,7 +161,7 @@ def main(): torch.cuda.synchronize() self.t0 = time.perf_counter() - def on_log(self, _args, state, control, logs=None, **kwargs): + def on_log(self, _args, state, control, logs = None, **kwargs): if logs is None: return if "loss" in logs: @@ -158,12 +175,12 @@ def main(): timings["step_wall"].append(time.perf_counter() - self.t0) trainer = GRPOTrainer( - model=model, - processing_class=tokenizer, - reward_funcs=reward_funcs, - args=training_args, - train_dataset=dataset, - callbacks=[StepTimer()], + model = model, + processing_class = tokenizer, + reward_funcs = reward_funcs, + args = training_args, + train_dataset = dataset, + callbacks = [StepTimer()], ) torch.cuda.reset_peak_memory_stats() @@ -187,7 +204,7 @@ def main(): "max_steps": args.max_steps, } with open(args.stats_path, "w") as f: - json.dump(stats, f, indent=2) + json.dump(stats, f, indent = 2) print(f"[naive] Wrote stats to {args.stats_path}") print(f"[naive] Total train wall: {t_train:.1f}s Peak mem: {peak:.2f} GB") diff --git a/scripts/benchmarks/qwen3_grpo_tpaged.py b/scripts/benchmarks/qwen3_grpo_tpaged.py index 321c8b1aed..13d6743427 100644 --- a/scripts/benchmarks/qwen3_grpo_tpaged.py +++ b/scripts/benchmarks/qwen3_grpo_tpaged.py @@ -31,10 +31,13 @@ sys.path.insert(0, str(HERE)) # `for_inference()` hooks, which a vanilla HF model does not. try: import vllm.sampling_params as _vllm_sp + if not hasattr(_vllm_sp, "GuidedDecodingParams"): + class _GuidedDecodingParamsShim: # pragma: no cover - used only if TRL asks def __init__(self, *a, **kw): pass + _vllm_sp.GuidedDecodingParams = _GuidedDecodingParamsShim except ImportError: pass @@ -44,6 +47,7 @@ from transformers import AutoModelForCausalLM, AutoTokenizer # noqa: E402 from peft import LoraConfig, get_peft_model # noqa: E402 import flash_attn_fa4_shim # noqa: E402 + flash_attn_fa4_shim.apply() from unsloth_grpo_common import ( # noqa: E402 @@ -56,31 +60,45 @@ from unsloth_grpo_common import ( # noqa: E402 def parse_args(): p = argparse.ArgumentParser() - 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=61) - 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("--attn_impl", default="sdpa", - help="Base attention impl to compose with paged. 'sdpa' or 'flash_attention_2'.") - p.add_argument("--max_batch_tokens", type=int, default=8192, - help="PagedAttentionCache.max_batch_tokens. Default upper bound is 256 which is far too small.") - p.add_argument("--num_blocks", type=int, default=8192, - help="PagedAttentionCache.num_blocks (block_size=32). 8192*32 tokens of KV capacity.") - p.add_argument("--output_dir", default="outputs/grpo_tpaged") - p.add_argument("--stats_path", default="logs/tpaged_stats.json") - p.add_argument("--persistent_cb", action="store_true", - help="Reuse one ContinuousBatchingManager across every training step instead " - "of letting TRL's generate_batch rebuild it (and the paged cache) each step.") + 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 = 61) + 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( + "--attn_impl", + default = "sdpa", + help = "Base attention impl to compose with paged. 'sdpa' or 'flash_attention_2'.", + ) + p.add_argument( + "--max_batch_tokens", + type = int, + default = 8192, + help = "PagedAttentionCache.max_batch_tokens. Default upper bound is 256 which is far too small.", + ) + p.add_argument( + "--num_blocks", + type = int, + default = 8192, + help = "PagedAttentionCache.num_blocks (block_size=32). 8192*32 tokens of KV capacity.", + ) + p.add_argument("--output_dir", default = "outputs/grpo_tpaged") + p.add_argument("--stats_path", default = "logs/tpaged_stats.json") + p.add_argument( + "--persistent_cb", + action = "store_true", + help = "Reuse one ContinuousBatchingManager across every training step instead " + "of letting TRL's generate_batch rebuild it (and the paged cache) each step.", + ) return p.parse_args() 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) # 1. Vanilla HF load (no Unsloth patches on the attention forward). tokenizer = AutoTokenizer.from_pretrained(args.model_name) @@ -88,24 +106,31 @@ def main(): 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, ) model.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", + 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", + bias = "none", + task_type = "CAUSAL_LM", ) model = get_peft_model(model, lora) try: - model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}) + model.gradient_checkpointing_enable( + gradient_checkpointing_kwargs = {"use_reentrant": False} + ) except TypeError: model.gradient_checkpointing_enable() model.enable_input_require_grads() @@ -113,21 +138,24 @@ def main(): apply_chat_template_to_tokenizer(tokenizer) # 2. Dataset + rewards (identical to the vLLM script). - 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"[tpaged] Max prompt length (p90): {maximum_length}") reward_funcs = build_reward_funcs(tokenizer) # 3. Build GRPOConfig with transformers continuous batching enabled. from trl import GRPOConfig, GRPOTrainer + 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, ) # transformers `TopKLogitsWarper` rejects -1. None skips the warper entirely. shared["top_k"] = None @@ -137,10 +165,10 @@ def main(): # `generation_kwargs`, which TRL forwards to `GenerationConfig`, which the # CB manager then reads off when sizing the paged cache. 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, }, @@ -160,7 +188,7 @@ def main(): torch.cuda.synchronize() self.t0 = time.perf_counter() - def on_log(self, _args, state, control, logs=None, **kwargs): + def on_log(self, _args, state, control, logs = None, **kwargs): if logs is None: return if "loss" in logs: @@ -174,21 +202,26 @@ def main(): timings["step_wall"].append(time.perf_counter() - self.t0) trainer = GRPOTrainer( - model=model, - processing_class=tokenizer, - reward_funcs=reward_funcs, - args=training_args, - train_dataset=dataset, - callbacks=[StepTimer()], + model = model, + processing_class = tokenizer, + reward_funcs = reward_funcs, + args = training_args, + train_dataset = dataset, + callbacks = [StepTimer()], ) if args.persistent_cb: # TRL constructs `self.generation_config` once in `__init__`; reuse # the same object so the persistent manager stays warm. from persistent_cb import install_for_model, teardown + # TRL generates against the unwrapped base model; attach the patch # directly to it so every rollout picks up the persistent manager. - 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) # PEFT's wrapper chains `generate_batch` through `base_model.model` via # __getattr__, so installing on `base` is enough for TRL's call path. @@ -200,6 +233,7 @@ def main(): finally: if args.persistent_cb: from persistent_cb import teardown + teardown(base) t_train = time.perf_counter() - t_start @@ -220,7 +254,7 @@ def main(): "persistent_cb": args.persistent_cb, } with open(args.stats_path, "w") as f: - json.dump(stats, f, indent=2) + json.dump(stats, f, indent = 2) print(f"[tpaged] Wrote stats to {args.stats_path}") print(f"[tpaged] Total train wall: {t_train:.1f}s Peak mem: {peak:.2f} GB")