[pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci
This commit is contained in:
pre-commit-ci[bot] 2026-04-19 14:45:46 +00:00
commit 16d80d7378
4 changed files with 216 additions and 169 deletions

View file

@ -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__":

View file

@ -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")

View file

@ -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")

View file

@ -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"</SOLUTION>[\s]{0,}"
+ "(?:" + re.escape(tokenizer.eos_token) + ")?"
r"</SOLUTION>[\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,
)