[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
This commit is contained in:
parent
07939bb025
commit
16d80d7378
4 changed files with 216 additions and 169 deletions
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue