[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 01:38:36 +00:00
commit 57a7aefd94
4 changed files with 215 additions and 144 deletions

View file

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

View file

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

View file

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

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