hard007ik commited on
Commit
1ecce94
·
1 Parent(s): caea08d

Drop max_prompt_length from GRPOConfig (removed in newer TRL)

Browse files
Files changed (1) hide show
  1. 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
- ap.add_argument("--max-prompt-length", type=int, default=2048)
 
 
 
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=args.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,