SavK1 Claude Sonnet 4.6 commited on
Commit
22d40fb
·
1 Parent(s): 1b68f93

fix(train_v3): use PMOpsGRPOTrainer to fix reward kwarg plumbing

Browse files

PatchFastRL 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>

Files changed (1) hide show
  1. 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": {"display_name": "Python 3", "language": "python", "name": "python3"},
6
- "language_info": {"name": "python", "version": "3.11.0"}
 
 
 
 
 
 
 
7
  },
8
  "cells": [
9
  {
@@ -34,7 +41,9 @@
34
  "cell_type": "markdown",
35
  "id": "cell-1-md",
36
  "metadata": {},
37
- "source": ["## 0. Install"]
 
 
38
  },
39
  {
40
  "cell_type": "code",
@@ -58,7 +67,9 @@
58
  "cell_type": "markdown",
59
  "id": "cell-2-md",
60
  "metadata": {},
61
- "source": ["## 1. Imports + GPU Config"]
 
 
62
  },
63
  {
64
  "cell_type": "code",
@@ -98,7 +109,9 @@
98
  "cell_type": "markdown",
99
  "id": "cell-3-md",
100
  "metadata": {},
101
- "source": ["## 2. Clone PM-Ops Repo"]
 
 
102
  },
103
  {
104
  "cell_type": "code",
@@ -130,7 +143,9 @@
130
  "cell_type": "markdown",
131
  "id": "cell-4-md",
132
  "metadata": {},
133
- "source": ["## 3. HuggingFace Login"]
 
 
134
  },
135
  {
136
  "cell_type": "code",
@@ -147,7 +162,9 @@
147
  "cell_type": "markdown",
148
  "id": "cell-5-md",
149
  "metadata": {},
150
- "source": ["## 4. Start Local PM-Ops Server"]
 
 
151
  },
152
  {
153
  "cell_type": "code",
@@ -181,7 +198,9 @@
181
  "cell_type": "markdown",
182
  "id": "cell-6-md",
183
  "metadata": {},
184
- "source": ["## 5. Verify Env"]
 
 
185
  },
186
  {
187
  "cell_type": "code",
@@ -206,7 +225,9 @@
206
  "cell_type": "markdown",
207
  "id": "cell-7-md",
208
  "metadata": {},
209
- "source": ["## 6. Load Model — Unsloth 4-bit + LoRA"]
 
 
210
  },
211
  {
212
  "cell_type": "code",
@@ -264,7 +285,9 @@
264
  "cell_type": "markdown",
265
  "id": "cell-9-md",
266
  "metadata": {},
267
- "source": ["## 7. Generate SFT Demonstration Dataset"]
 
 
268
  },
269
  {
270
  "cell_type": "code",
@@ -344,7 +367,9 @@
344
  "cell_type": "markdown",
345
  "id": "cell-10-md",
346
  "metadata": {},
347
- "source": ["## 8. SFT Training (~15 min on A100)"]
 
 
348
  },
349
  {
350
  "cell_type": "code",
@@ -392,7 +417,9 @@
392
  "cell_type": "markdown",
393
  "id": "cell-11-md",
394
  "metadata": {},
395
- "source": ["## 9. Verify SFT Output — Model Must Output Valid JSON"]
 
 
396
  },
397
  {
398
  "cell_type": "code",
@@ -459,7 +486,9 @@
459
  "cell_type": "markdown",
460
  "id": "cell-13-md",
461
  "metadata": {},
462
- "source": ["## 10. Generate GRPO Training Dataset"]
 
 
463
  },
464
  {
465
  "cell_type": "code",
@@ -480,7 +509,9 @@
480
  "cell_type": "markdown",
481
  "id": "cell-14-md",
482
  "metadata": {},
483
- "source": ["## 11. GRPO Rollout + Reward Functions"]
 
 
484
  },
485
  {
486
  "cell_type": "code",
@@ -593,7 +624,9 @@
593
  "cell_type": "markdown",
594
  "id": "cell-15-md",
595
  "metadata": {},
596
- "source": ["## 12. GRPO Config + Trainer"]
 
 
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": ["## 14. Save Model"]
 
 
684
  },
685
  {
686
  "cell_type": "code",
@@ -702,7 +696,9 @@
702
  "cell_type": "markdown",
703
  "id": "cell-18-md",
704
  "metadata": {},
705
- "source": ["## 15. Evaluate: Baseline vs Trained"]
 
 
706
  },
707
  {
708
  "cell_type": "code",
@@ -795,7 +791,9 @@
795
  "cell_type": "markdown",
796
  "id": "cell-19-md",
797
  "metadata": {},
798
- "source": ["## 16. Plot Results"]
 
 
799
  },
800
  {
801
  "cell_type": "code",
@@ -837,7 +835,9 @@
837
  "cell_type": "markdown",
838
  "id": "cell-20-md",
839
  "metadata": {},
840
- "source": ["## 17. Teardown"]
 
 
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
+ }