[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-20 13:54:21 +00:00
commit c31533fc02
4 changed files with 471 additions and 326 deletions

View file

@ -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,31 +257,37 @@ 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)
# PEFT saves with keys like `base_model.model.model.layers.0.self_attn.q_proj.lora_A.weight`.
@ -283,22 +298,28 @@ def run_unsloth_fi_false(args):
with torch.no_grad():
for name, tensor in loaded_tensors.items():
# Try direct + strip `base_model.model.` prefix variants.
candidates = [name, name.replace("base_model.model.", ""),
"base_model.model." + name]
candidates = [
name,
name.replace("base_model.model.", ""),
"base_model.model." + name,
]
for cand in candidates:
# PEFT sometimes inserts `.default.` between module and lora_A.
variants = [cand, cand.replace(".default.", ".")]
for v in variants:
# own_state keys typically have `.default.weight` suffix
for own_name, own in own_state.items():
if own_name.endswith(v.split("base_model.model.")[-1]) \
or v.endswith(own_name.split("base_model.model.")[-1]):
if own_name.endswith(
v.split("base_model.model.")[-1]
) or v.endswith(own_name.split("base_model.model.")[-1]):
if own.shape == tensor.shape:
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)
@ -306,28 +327,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
@ -348,7 +370,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
@ -356,8 +380,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",
@ -376,30 +403,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":
@ -417,8 +449,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)

View file

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

View file

@ -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 = "<start_working_out>"
@ -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("<think>", "").replace("</think>", "").strip()
thoughts = (
x["generated_solution"]
.replace("<think>", "")
.replace("</think>", "")
.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"</SOLUTION>[\s]{0,}" + "(?:" + re.escape(tokenizer.eos_token) + ")?"
solution_end_regex = (
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):
@ -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:

View file

@ -42,28 +42,30 @@ 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)
return p.parse_args()
@ -71,10 +73,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():
@ -82,11 +92,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,
)
@ -96,10 +107,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
@ -107,22 +121,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
@ -138,21 +164,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()
@ -162,38 +195,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)
@ -201,12 +242,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
@ -217,18 +258,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":
@ -236,16 +283,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,
},
@ -253,27 +300,30 @@ 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
@ -284,6 +334,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
@ -323,10 +374,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)