Spaces:
Sleeping
Sleeping
fix(train_v3): use PMOpsGRPOTrainer to fix reward kwarg plumbing
Browse filesPatchFastRL strips custom rollout keys before calling reward_func,
so the 'reward' key never arrived. PMOpsGRPOTrainer overrides
_calculate_rewards to inject pre-computed rewards from the rollout
batch directly, bypassing the broken kwargs path entirely.
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
- training/train_v3.ipynb +62 -62
training/train_v3.ipynb
CHANGED
|
@@ -2,8 +2,15 @@
|
|
| 2 |
"nbformat": 4,
|
| 3 |
"nbformat_minor": 5,
|
| 4 |
"metadata": {
|
| 5 |
-
"kernelspec": {
|
| 6 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 7 |
},
|
| 8 |
"cells": [
|
| 9 |
{
|
|
@@ -34,7 +41,9 @@
|
|
| 34 |
"cell_type": "markdown",
|
| 35 |
"id": "cell-1-md",
|
| 36 |
"metadata": {},
|
| 37 |
-
"source": [
|
|
|
|
|
|
|
| 38 |
},
|
| 39 |
{
|
| 40 |
"cell_type": "code",
|
|
@@ -58,7 +67,9 @@
|
|
| 58 |
"cell_type": "markdown",
|
| 59 |
"id": "cell-2-md",
|
| 60 |
"metadata": {},
|
| 61 |
-
"source": [
|
|
|
|
|
|
|
| 62 |
},
|
| 63 |
{
|
| 64 |
"cell_type": "code",
|
|
@@ -98,7 +109,9 @@
|
|
| 98 |
"cell_type": "markdown",
|
| 99 |
"id": "cell-3-md",
|
| 100 |
"metadata": {},
|
| 101 |
-
"source": [
|
|
|
|
|
|
|
| 102 |
},
|
| 103 |
{
|
| 104 |
"cell_type": "code",
|
|
@@ -130,7 +143,9 @@
|
|
| 130 |
"cell_type": "markdown",
|
| 131 |
"id": "cell-4-md",
|
| 132 |
"metadata": {},
|
| 133 |
-
"source": [
|
|
|
|
|
|
|
| 134 |
},
|
| 135 |
{
|
| 136 |
"cell_type": "code",
|
|
@@ -147,7 +162,9 @@
|
|
| 147 |
"cell_type": "markdown",
|
| 148 |
"id": "cell-5-md",
|
| 149 |
"metadata": {},
|
| 150 |
-
"source": [
|
|
|
|
|
|
|
| 151 |
},
|
| 152 |
{
|
| 153 |
"cell_type": "code",
|
|
@@ -181,7 +198,9 @@
|
|
| 181 |
"cell_type": "markdown",
|
| 182 |
"id": "cell-6-md",
|
| 183 |
"metadata": {},
|
| 184 |
-
"source": [
|
|
|
|
|
|
|
| 185 |
},
|
| 186 |
{
|
| 187 |
"cell_type": "code",
|
|
@@ -206,7 +225,9 @@
|
|
| 206 |
"cell_type": "markdown",
|
| 207 |
"id": "cell-7-md",
|
| 208 |
"metadata": {},
|
| 209 |
-
"source": [
|
|
|
|
|
|
|
| 210 |
},
|
| 211 |
{
|
| 212 |
"cell_type": "code",
|
|
@@ -264,7 +285,9 @@
|
|
| 264 |
"cell_type": "markdown",
|
| 265 |
"id": "cell-9-md",
|
| 266 |
"metadata": {},
|
| 267 |
-
"source": [
|
|
|
|
|
|
|
| 268 |
},
|
| 269 |
{
|
| 270 |
"cell_type": "code",
|
|
@@ -344,7 +367,9 @@
|
|
| 344 |
"cell_type": "markdown",
|
| 345 |
"id": "cell-10-md",
|
| 346 |
"metadata": {},
|
| 347 |
-
"source": [
|
|
|
|
|
|
|
| 348 |
},
|
| 349 |
{
|
| 350 |
"cell_type": "code",
|
|
@@ -392,7 +417,9 @@
|
|
| 392 |
"cell_type": "markdown",
|
| 393 |
"id": "cell-11-md",
|
| 394 |
"metadata": {},
|
| 395 |
-
"source": [
|
|
|
|
|
|
|
| 396 |
},
|
| 397 |
{
|
| 398 |
"cell_type": "code",
|
|
@@ -459,7 +486,9 @@
|
|
| 459 |
"cell_type": "markdown",
|
| 460 |
"id": "cell-13-md",
|
| 461 |
"metadata": {},
|
| 462 |
-
"source": [
|
|
|
|
|
|
|
| 463 |
},
|
| 464 |
{
|
| 465 |
"cell_type": "code",
|
|
@@ -480,7 +509,9 @@
|
|
| 480 |
"cell_type": "markdown",
|
| 481 |
"id": "cell-14-md",
|
| 482 |
"metadata": {},
|
| 483 |
-
"source": [
|
|
|
|
|
|
|
| 484 |
},
|
| 485 |
{
|
| 486 |
"cell_type": "code",
|
|
@@ -593,7 +624,9 @@
|
|
| 593 |
"cell_type": "markdown",
|
| 594 |
"id": "cell-15-md",
|
| 595 |
"metadata": {},
|
| 596 |
-
"source": [
|
|
|
|
|
|
|
| 597 |
},
|
| 598 |
{
|
| 599 |
"cell_type": "code",
|
|
@@ -601,48 +634,7 @@
|
|
| 601 |
"metadata": {},
|
| 602 |
"outputs": [],
|
| 603 |
"execution_count": null,
|
| 604 |
-
"source":
|
| 605 |
-
"from trl import GRPOConfig, GRPOTrainer\n",
|
| 606 |
-
"\n",
|
| 607 |
-
"OUTPUT_DIR = 'pm-ops-grpo-Qwen3-1.7B-triage-v3'\n",
|
| 608 |
-
"HF_REPO_ID = f'Saurav1/{OUTPUT_DIR}'\n",
|
| 609 |
-
"\n",
|
| 610 |
-
"grpo_cfg = GRPOConfig(\n",
|
| 611 |
-
" # Training\n",
|
| 612 |
-
" num_train_epochs = 2,\n",
|
| 613 |
-
" learning_rate = 1e-6, # lower LR: model has SFT init, don't overwrite it\n",
|
| 614 |
-
" gradient_accumulation_steps = GRAD_ACCUM,\n",
|
| 615 |
-
" per_device_train_batch_size = 1,\n",
|
| 616 |
-
" warmup_steps = 5,\n",
|
| 617 |
-
" num_generations = NUM_GEN,\n",
|
| 618 |
-
" # Sequence lengths\n",
|
| 619 |
-
" max_completion_length = MAX_COMP_LEN,\n",
|
| 620 |
-
" max_prompt_length = 4096,\n",
|
| 621 |
-
" # Unsloth vLLM for fast generation\n",
|
| 622 |
-
" use_vllm = True,\n",
|
| 623 |
-
" # Output\n",
|
| 624 |
-
" output_dir = OUTPUT_DIR,\n",
|
| 625 |
-
" report_to = 'trackio',\n",
|
| 626 |
-
" trackio_space_id = OUTPUT_DIR,\n",
|
| 627 |
-
" logging_steps = 1,\n",
|
| 628 |
-
" save_steps = 20,\n",
|
| 629 |
-
" gradient_checkpointing = False, # Unsloth handles this\n",
|
| 630 |
-
")\n",
|
| 631 |
-
"\n",
|
| 632 |
-
"eff_batch = grpo_cfg.per_device_train_batch_size * GRAD_ACCUM\n",
|
| 633 |
-
"total_steps = len(grpo_dataset) * NUM_GEN * grpo_cfg.num_train_epochs // eff_batch\n",
|
| 634 |
-
"print(f'GRPO: {len(grpo_dataset)} eps x {NUM_GEN} gen x {grpo_cfg.num_train_epochs} epochs -> ~{total_steps} steps')\n",
|
| 635 |
-
"\n",
|
| 636 |
-
"trainer = GRPOTrainer(\n",
|
| 637 |
-
" model = model,\n",
|
| 638 |
-
" processing_class = tokenizer,\n",
|
| 639 |
-
" reward_funcs = grpo_reward_func,\n",
|
| 640 |
-
" train_dataset = grpo_dataset,\n",
|
| 641 |
-
" args = grpo_cfg,\n",
|
| 642 |
-
" rollout_func = grpo_rollout_func,\n",
|
| 643 |
-
")\n",
|
| 644 |
-
"print('GRPOTrainer ready')"
|
| 645 |
-
]
|
| 646 |
},
|
| 647 |
{
|
| 648 |
"cell_type": "markdown",
|
|
@@ -680,7 +672,9 @@
|
|
| 680 |
"cell_type": "markdown",
|
| 681 |
"id": "cell-17-md",
|
| 682 |
"metadata": {},
|
| 683 |
-
"source": [
|
|
|
|
|
|
|
| 684 |
},
|
| 685 |
{
|
| 686 |
"cell_type": "code",
|
|
@@ -702,7 +696,9 @@
|
|
| 702 |
"cell_type": "markdown",
|
| 703 |
"id": "cell-18-md",
|
| 704 |
"metadata": {},
|
| 705 |
-
"source": [
|
|
|
|
|
|
|
| 706 |
},
|
| 707 |
{
|
| 708 |
"cell_type": "code",
|
|
@@ -795,7 +791,9 @@
|
|
| 795 |
"cell_type": "markdown",
|
| 796 |
"id": "cell-19-md",
|
| 797 |
"metadata": {},
|
| 798 |
-
"source": [
|
|
|
|
|
|
|
| 799 |
},
|
| 800 |
{
|
| 801 |
"cell_type": "code",
|
|
@@ -837,7 +835,9 @@
|
|
| 837 |
"cell_type": "markdown",
|
| 838 |
"id": "cell-20-md",
|
| 839 |
"metadata": {},
|
| 840 |
-
"source": [
|
|
|
|
|
|
|
| 841 |
},
|
| 842 |
{
|
| 843 |
"cell_type": "code",
|
|
@@ -851,4 +851,4 @@
|
|
| 851 |
]
|
| 852 |
}
|
| 853 |
]
|
| 854 |
-
}
|
|
|
|
| 2 |
"nbformat": 4,
|
| 3 |
"nbformat_minor": 5,
|
| 4 |
"metadata": {
|
| 5 |
+
"kernelspec": {
|
| 6 |
+
"display_name": "Python 3",
|
| 7 |
+
"language": "python",
|
| 8 |
+
"name": "python3"
|
| 9 |
+
},
|
| 10 |
+
"language_info": {
|
| 11 |
+
"name": "python",
|
| 12 |
+
"version": "3.11.0"
|
| 13 |
+
}
|
| 14 |
},
|
| 15 |
"cells": [
|
| 16 |
{
|
|
|
|
| 41 |
"cell_type": "markdown",
|
| 42 |
"id": "cell-1-md",
|
| 43 |
"metadata": {},
|
| 44 |
+
"source": [
|
| 45 |
+
"## 0. Install"
|
| 46 |
+
]
|
| 47 |
},
|
| 48 |
{
|
| 49 |
"cell_type": "code",
|
|
|
|
| 67 |
"cell_type": "markdown",
|
| 68 |
"id": "cell-2-md",
|
| 69 |
"metadata": {},
|
| 70 |
+
"source": [
|
| 71 |
+
"## 1. Imports + GPU Config"
|
| 72 |
+
]
|
| 73 |
},
|
| 74 |
{
|
| 75 |
"cell_type": "code",
|
|
|
|
| 109 |
"cell_type": "markdown",
|
| 110 |
"id": "cell-3-md",
|
| 111 |
"metadata": {},
|
| 112 |
+
"source": [
|
| 113 |
+
"## 2. Clone PM-Ops Repo"
|
| 114 |
+
]
|
| 115 |
},
|
| 116 |
{
|
| 117 |
"cell_type": "code",
|
|
|
|
| 143 |
"cell_type": "markdown",
|
| 144 |
"id": "cell-4-md",
|
| 145 |
"metadata": {},
|
| 146 |
+
"source": [
|
| 147 |
+
"## 3. HuggingFace Login"
|
| 148 |
+
]
|
| 149 |
},
|
| 150 |
{
|
| 151 |
"cell_type": "code",
|
|
|
|
| 162 |
"cell_type": "markdown",
|
| 163 |
"id": "cell-5-md",
|
| 164 |
"metadata": {},
|
| 165 |
+
"source": [
|
| 166 |
+
"## 4. Start Local PM-Ops Server"
|
| 167 |
+
]
|
| 168 |
},
|
| 169 |
{
|
| 170 |
"cell_type": "code",
|
|
|
|
| 198 |
"cell_type": "markdown",
|
| 199 |
"id": "cell-6-md",
|
| 200 |
"metadata": {},
|
| 201 |
+
"source": [
|
| 202 |
+
"## 5. Verify Env"
|
| 203 |
+
]
|
| 204 |
},
|
| 205 |
{
|
| 206 |
"cell_type": "code",
|
|
|
|
| 225 |
"cell_type": "markdown",
|
| 226 |
"id": "cell-7-md",
|
| 227 |
"metadata": {},
|
| 228 |
+
"source": [
|
| 229 |
+
"## 6. Load Model — Unsloth 4-bit + LoRA"
|
| 230 |
+
]
|
| 231 |
},
|
| 232 |
{
|
| 233 |
"cell_type": "code",
|
|
|
|
| 285 |
"cell_type": "markdown",
|
| 286 |
"id": "cell-9-md",
|
| 287 |
"metadata": {},
|
| 288 |
+
"source": [
|
| 289 |
+
"## 7. Generate SFT Demonstration Dataset"
|
| 290 |
+
]
|
| 291 |
},
|
| 292 |
{
|
| 293 |
"cell_type": "code",
|
|
|
|
| 367 |
"cell_type": "markdown",
|
| 368 |
"id": "cell-10-md",
|
| 369 |
"metadata": {},
|
| 370 |
+
"source": [
|
| 371 |
+
"## 8. SFT Training (~15 min on A100)"
|
| 372 |
+
]
|
| 373 |
},
|
| 374 |
{
|
| 375 |
"cell_type": "code",
|
|
|
|
| 417 |
"cell_type": "markdown",
|
| 418 |
"id": "cell-11-md",
|
| 419 |
"metadata": {},
|
| 420 |
+
"source": [
|
| 421 |
+
"## 9. Verify SFT Output — Model Must Output Valid JSON"
|
| 422 |
+
]
|
| 423 |
},
|
| 424 |
{
|
| 425 |
"cell_type": "code",
|
|
|
|
| 486 |
"cell_type": "markdown",
|
| 487 |
"id": "cell-13-md",
|
| 488 |
"metadata": {},
|
| 489 |
+
"source": [
|
| 490 |
+
"## 10. Generate GRPO Training Dataset"
|
| 491 |
+
]
|
| 492 |
},
|
| 493 |
{
|
| 494 |
"cell_type": "code",
|
|
|
|
| 509 |
"cell_type": "markdown",
|
| 510 |
"id": "cell-14-md",
|
| 511 |
"metadata": {},
|
| 512 |
+
"source": [
|
| 513 |
+
"## 11. GRPO Rollout + Reward Functions"
|
| 514 |
+
]
|
| 515 |
},
|
| 516 |
{
|
| 517 |
"cell_type": "code",
|
|
|
|
| 624 |
"cell_type": "markdown",
|
| 625 |
"id": "cell-15-md",
|
| 626 |
"metadata": {},
|
| 627 |
+
"source": [
|
| 628 |
+
"## 12. GRPO Config + Trainer"
|
| 629 |
+
]
|
| 630 |
},
|
| 631 |
{
|
| 632 |
"cell_type": "code",
|
|
|
|
| 634 |
"metadata": {},
|
| 635 |
"outputs": [],
|
| 636 |
"execution_count": null,
|
| 637 |
+
"source": "from trl import GRPOConfig\nfrom training.pm_ops_trainer import PMOpsGRPOTrainer\n\nOUTPUT_DIR = 'pm-ops-grpo-Qwen3-1.7B-triage-v3'\nHF_REPO_ID = f'Saurav1/{OUTPUT_DIR}'\n\ngrpo_cfg = GRPOConfig(\n # Training\n num_train_epochs = 2,\n learning_rate = 1e-6, # lower LR: model has SFT init, don't overwrite it\n gradient_accumulation_steps = GRAD_ACCUM,\n per_device_train_batch_size = 1,\n warmup_steps = 5,\n num_generations = NUM_GEN,\n # Sequence lengths\n max_completion_length = MAX_COMP_LEN,\n max_prompt_length = 4096,\n # Unsloth vLLM for fast generation\n use_vllm = True,\n # Output\n output_dir = OUTPUT_DIR,\n report_to = 'trackio',\n trackio_space_id = OUTPUT_DIR,\n logging_steps = 1,\n save_steps = 20,\n gradient_checkpointing = False, # Unsloth handles this\n)\n\neff_batch = grpo_cfg.per_device_train_batch_size * GRAD_ACCUM\ntotal_steps = len(grpo_dataset) * NUM_GEN * grpo_cfg.num_train_epochs // eff_batch\nprint(f'GRPO: {len(grpo_dataset)} eps x {NUM_GEN} gen x {grpo_cfg.num_train_epochs} epochs -> ~{total_steps} steps')\n\n# PMOpsGRPOTrainer overrides _calculate_rewards to read the pre-computed\n# 'reward' key directly from the rollout batch — bypasses the broken\n# PatchFastRL kwargs plumbing that strips custom rollout keys.\ntrainer = PMOpsGRPOTrainer(\n model = model,\n processing_class = tokenizer,\n reward_funcs = grpo_reward_func, # kept as fallback only\n train_dataset = grpo_dataset,\n args = grpo_cfg,\n rollout_func = grpo_rollout_func,\n)\nprint('PMOpsGRPOTrainer ready')"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 638 |
},
|
| 639 |
{
|
| 640 |
"cell_type": "markdown",
|
|
|
|
| 672 |
"cell_type": "markdown",
|
| 673 |
"id": "cell-17-md",
|
| 674 |
"metadata": {},
|
| 675 |
+
"source": [
|
| 676 |
+
"## 14. Save Model"
|
| 677 |
+
]
|
| 678 |
},
|
| 679 |
{
|
| 680 |
"cell_type": "code",
|
|
|
|
| 696 |
"cell_type": "markdown",
|
| 697 |
"id": "cell-18-md",
|
| 698 |
"metadata": {},
|
| 699 |
+
"source": [
|
| 700 |
+
"## 15. Evaluate: Baseline vs Trained"
|
| 701 |
+
]
|
| 702 |
},
|
| 703 |
{
|
| 704 |
"cell_type": "code",
|
|
|
|
| 791 |
"cell_type": "markdown",
|
| 792 |
"id": "cell-19-md",
|
| 793 |
"metadata": {},
|
| 794 |
+
"source": [
|
| 795 |
+
"## 16. Plot Results"
|
| 796 |
+
]
|
| 797 |
},
|
| 798 |
{
|
| 799 |
"cell_type": "code",
|
|
|
|
| 835 |
"cell_type": "markdown",
|
| 836 |
"id": "cell-20-md",
|
| 837 |
"metadata": {},
|
| 838 |
+
"source": [
|
| 839 |
+
"## 17. Teardown"
|
| 840 |
+
]
|
| 841 |
},
|
| 842 |
{
|
| 843 |
"cell_type": "code",
|
|
|
|
| 851 |
]
|
| 852 |
}
|
| 853 |
]
|
| 854 |
+
}
|