From 16d80d73786ed264b562eb62c94e623bef12dce2 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sun, 19 Apr 2026 14:45:46 +0000 Subject: [PATCH] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- scripts/benchmarks/cb_vs_vllm_generation.py | 98 +++++++++-------- scripts/benchmarks/qwen3_grpo_tpaged.py | 116 ++++++++++++-------- scripts/benchmarks/qwen3_grpo_vllm.py | 109 +++++++++--------- scripts/benchmarks/unsloth_grpo_common.py | 62 ++++++----- 4 files changed, 216 insertions(+), 169 deletions(-) diff --git a/scripts/benchmarks/cb_vs_vllm_generation.py b/scripts/benchmarks/cb_vs_vllm_generation.py index d949ecc111..b0b6e26966 100644 --- a/scripts/benchmarks/cb_vs_vllm_generation.py +++ b/scripts/benchmarks/cb_vs_vllm_generation.py @@ -36,8 +36,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}, @@ -46,11 +46,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 @@ -58,30 +58,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. @@ -91,7 +97,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) @@ -121,22 +129,22 @@ 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() 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 @@ -145,7 +153,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) @@ -156,7 +166,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) @@ -180,23 +190,23 @@ 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("--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("--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": @@ -206,8 +216,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)) if __name__ == "__main__": diff --git a/scripts/benchmarks/qwen3_grpo_tpaged.py b/scripts/benchmarks/qwen3_grpo_tpaged.py index 0efb028f51..ca128d2401 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 @@ -53,28 +56,39 @@ 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("--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") 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) @@ -82,24 +96,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() @@ -107,21 +128,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 @@ -131,10 +155,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, }, @@ -154,7 +178,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: @@ -168,12 +192,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() @@ -197,7 +221,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"[tpaged] Wrote stats to {args.stats_path}") print(f"[tpaged] Total train wall: {t_train:.1f}s Peak mem: {peak:.2f} GB") diff --git a/scripts/benchmarks/qwen3_grpo_vllm.py b/scripts/benchmarks/qwen3_grpo_vllm.py index 9673386aaa..4dd8073d74 100644 --- a/scripts/benchmarks/qwen3_grpo_vllm.py +++ b/scripts/benchmarks/qwen3_grpo_vllm.py @@ -35,79 +35,88 @@ 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("--gpu_memory_utilization", type=float, default=0.8) - p.add_argument("--output_dir", default="outputs/grpo_vllm") - p.add_argument("--stats_path", default="logs/vllm_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 = 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("--gpu_memory_utilization", type = float, default = 0.8) + p.add_argument("--output_dir", default = "outputs/grpo_vllm") + p.add_argument("--stats_path", default = "logs/vllm_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) # 1. Load model with vLLM fast inference enabled. 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, ) apply_chat_template_to_tokenizer(tokenizer) # 2. Dataset + rewards. - 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"[vllm] Max prompt length (p90): {maximum_length}") reward_funcs = build_reward_funcs(tokenizer) # 3. vLLM sampling params match the notebook. from vllm import SamplingParams + vllm_sampling_params = SamplingParams( - min_p=0.1, - top_p=1.0, - top_k=-1, - seed=3407, - stop=[tokenizer.eos_token], - include_stop_str_in_output=True, + min_p = 0.1, + top_p = 1.0, + top_k = -1, + seed = 3407, + stop = [tokenizer.eos_token], + include_stop_str_in_output = True, ) # 4. Build GRPOConfig. 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, ) training_args = GRPOConfig( - use_vllm=True, - vllm_mode="colocate", - vllm_sampling_params=vllm_sampling_params, - vllm_gpu_memory_utilization=args.gpu_memory_utilization, + use_vllm = True, + vllm_mode = "colocate", + vllm_sampling_params = vllm_sampling_params, + vllm_gpu_memory_utilization = args.gpu_memory_utilization, **shared, ) @@ -124,7 +133,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: @@ -138,12 +147,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() @@ -166,7 +175,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"[vllm] Wrote stats to {args.stats_path}") print(f"[vllm] Total train wall: {t_train:.1f}s Peak mem: {peak:.2f} GB") diff --git a/scripts/benchmarks/unsloth_grpo_common.py b/scripts/benchmarks/unsloth_grpo_common.py index d4cc7258ae..ab9921f540 100644 --- a/scripts/benchmarks/unsloth_grpo_common.py +++ b/scripts/benchmarks/unsloth_grpo_common.py @@ -63,7 +63,7 @@ def build_dataset(tokenizer, *, max_seq_length: int = 2048): """Build the DAPO-Math-17k GRPO dataset with the prompt formatting from the notebook. Returns `(dataset, maximum_prompt_length)`. """ - ds = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split="train") + ds = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split = "train") def _map_row(x): return { @@ -81,12 +81,12 @@ def build_dataset(tokenizer, *, max_seq_length: int = 2048): return { "tokens": tokenizer.apply_chat_template( batch["prompt"], - add_generation_prompt=True, - tokenize=True, + add_generation_prompt = True, + tokenize = True, ) } - tokenized = ds.map(_tokenize, batched=True) + tokenized = ds.map(_tokenize, batched = True) tokenized = tokenized.map(lambda x: {"L": len(x["tokens"])}) lengths = np.array(tokenized["L"]) maximum_length = int(np.quantile(lengths, 0.9)) @@ -97,18 +97,17 @@ def build_dataset(tokenizer, *, max_seq_length: int = 2048): def build_reward_funcs(tokenizer): """Return the 4 reward functions used in the notebook, wired to `tokenizer`.""" solution_end_regex = ( - r"[\s]{0,}" - + "(?:" + re.escape(tokenizer.eos_token) + ")?" + 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): @@ -193,7 +192,12 @@ def build_reward_funcs(tokenizer): scores.append(0.0) return scores - return [match_format_exactly, match_format_approximately, check_answer, check_numbers] + return [ + match_format_exactly, + match_format_approximately, + check_answer, + check_numbers, + ] def build_grpo_kwargs( @@ -215,24 +219,24 @@ def build_grpo_kwargs( max_completion_length = max_seq_length - max_prompt_length return dict( - temperature=1.0, - top_p=1.0, - top_k=-1, - min_p=0.1, - 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=per_device_train_batch_size, - gradient_accumulation_steps=gradient_accumulation_steps, - num_generations=num_generations, - max_prompt_length=max_prompt_length, - max_completion_length=max_completion_length, - max_steps=max_steps, - save_steps=max_steps, - report_to="none", - output_dir=output_dir, - seed=3407, + temperature = 1.0, + top_p = 1.0, + top_k = -1, + min_p = 0.1, + 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 = per_device_train_batch_size, + gradient_accumulation_steps = gradient_accumulation_steps, + num_generations = num_generations, + max_prompt_length = max_prompt_length, + max_completion_length = max_completion_length, + max_steps = max_steps, + save_steps = max_steps, + report_to = "none", + output_dir = output_dir, + seed = 3407, )