Drop max_prompt_length from GRPOConfig (removed in newer TRL)
Browse files- train_jewelry_grpo.py +5 -3
train_jewelry_grpo.py
CHANGED
|
@@ -119,7 +119,10 @@ def main() -> None:
|
|
| 119 |
ap.add_argument("--per-device-batch", type=int, default=1)
|
| 120 |
ap.add_argument("--grad-accum", type=int, default=32)
|
| 121 |
ap.add_argument("--max-completion-length", type=int, default=64)
|
| 122 |
-
|
|
|
|
|
|
|
|
|
|
| 123 |
ap.add_argument("--max-turns", type=int, default=15)
|
| 124 |
ap.add_argument("--lr", type=float, default=5e-6)
|
| 125 |
ap.add_argument("--warmup-steps", type=int, default=10)
|
|
@@ -251,7 +254,7 @@ def main() -> None:
|
|
| 251 |
warmup_steps=args.warmup_steps,
|
| 252 |
num_generations=args.num_generations,
|
| 253 |
max_completion_length=args.max_completion_length,
|
| 254 |
-
max_prompt_length
|
| 255 |
# vLLM is the canonical generation backend on GPU; turn off on CPU smoke.
|
| 256 |
use_vllm=not use_cpu,
|
| 257 |
vllm_mode="colocate" if not use_cpu else None,
|
|
@@ -291,7 +294,6 @@ def main() -> None:
|
|
| 291 |
"per_device_batch": args.per_device_batch,
|
| 292 |
"grad_accum": args.grad_accum,
|
| 293 |
"max_completion_length": args.max_completion_length,
|
| 294 |
-
"max_prompt_length": args.max_prompt_length,
|
| 295 |
"max_turns": args.max_turns,
|
| 296 |
"lr": args.lr,
|
| 297 |
"warmup_steps": args.warmup_steps,
|
|
|
|
| 119 |
ap.add_argument("--per-device-batch", type=int, default=1)
|
| 120 |
ap.add_argument("--grad-accum", type=int, default=32)
|
| 121 |
ap.add_argument("--max-completion-length", type=int, default=64)
|
| 122 |
+
# NOTE: --max-prompt-length intentionally removed. Recent TRL versions
|
| 123 |
+
# (>=0.20-ish) dropped `max_prompt_length` from GRPOConfig and now infer
|
| 124 |
+
# it from the tokenizer / model context window. Passing it back errors
|
| 125 |
+
# with `TypeError: unexpected keyword argument 'max_prompt_length'`.
|
| 126 |
ap.add_argument("--max-turns", type=int, default=15)
|
| 127 |
ap.add_argument("--lr", type=float, default=5e-6)
|
| 128 |
ap.add_argument("--warmup-steps", type=int, default=10)
|
|
|
|
| 254 |
warmup_steps=args.warmup_steps,
|
| 255 |
num_generations=args.num_generations,
|
| 256 |
max_completion_length=args.max_completion_length,
|
| 257 |
+
# `max_prompt_length` removed: not accepted by recent TRL GRPOConfig.
|
| 258 |
# vLLM is the canonical generation backend on GPU; turn off on CPU smoke.
|
| 259 |
use_vllm=not use_cpu,
|
| 260 |
vllm_mode="colocate" if not use_cpu else None,
|
|
|
|
| 294 |
"per_device_batch": args.per_device_batch,
|
| 295 |
"grad_accum": args.grad_accum,
|
| 296 |
"max_completion_length": args.max_completion_length,
|
|
|
|
| 297 |
"max_turns": args.max_turns,
|
| 298 |
"lr": args.lr,
|
| 299 |
"warmup_steps": args.warmup_steps,
|