piyush-mk commited on
Commit
59ed7f3
·
verified ·
1 Parent(s): a674764

v6: GRPO with --resume-adapter to start from SFT checkpoint

Browse files
Files changed (1) hide show
  1. training/train_grpo.py +21 -1
training/train_grpo.py CHANGED
@@ -186,6 +186,7 @@ class TrainConfig:
186
  format_warmup_lr: float = 5e-5
187
  save_format_warmup_checkpoint: bool = True
188
  format_warmup_model_id: Optional[str] = os.environ.get("FORMAT_WARMUP_MODEL_ID")
 
189
 
190
  # LoRA
191
  lora_r: int = 16
@@ -672,8 +673,24 @@ def train(cfg: TrainConfig) -> None:
672
  fig.savefig(artifact_dir / "holdout_eval_curve.png", dpi=160)
673
  plt.close(fig)
674
 
675
- # ----- Format warm-start ---------------------------------------------------
676
  global_step = 0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
677
  if cfg.format_warmup:
678
  print("\n=== format warm-start (JSON action behavior) ===", flush=True)
679
  warmup_metrics = run_format_warmup(
@@ -914,6 +931,7 @@ def _parse_args() -> TrainConfig:
914
  p.add_argument("--no-save-format-warmup", action="store_true")
915
  p.add_argument("--format-warmup-model-id", default=None)
916
  p.add_argument("--sample-temperature", type=float, default=None)
 
917
  args = p.parse_args()
918
 
919
  cfg = TrainConfig()
@@ -947,6 +965,8 @@ def _parse_args() -> TrainConfig:
947
  cfg.format_warmup_model_id = args.format_warmup_model_id
948
  if args.sample_temperature is not None:
949
  cfg.sample_temperature = args.sample_temperature
 
 
950
  if args.no_push:
951
  cfg.push_to_hub = False
952
  if args.no_4bit:
 
186
  format_warmup_lr: float = 5e-5
187
  save_format_warmup_checkpoint: bool = True
188
  format_warmup_model_id: Optional[str] = os.environ.get("FORMAT_WARMUP_MODEL_ID")
189
+ resume_adapter: Optional[str] = os.environ.get("RESUME_ADAPTER")
190
 
191
  # LoRA
192
  lora_r: int = 16
 
673
  fig.savefig(artifact_dir / "holdout_eval_curve.png", dpi=160)
674
  plt.close(fig)
675
 
676
+ # ----- Resume from existing adapter or format warm-start ------------------
677
  global_step = 0
678
+ if cfg.resume_adapter:
679
+ print(f"\n=== resuming from adapter: {cfg.resume_adapter} ===", flush=True)
680
+ from peft import set_peft_model_state_dict
681
+ from safetensors.torch import load_file
682
+ from huggingface_hub import hf_hub_download
683
+ try:
684
+ adapter_path = hf_hub_download(
685
+ cfg.resume_adapter, "adapter_model.safetensors", token=_hf_token()
686
+ )
687
+ adapter_weights = load_file(adapter_path)
688
+ set_peft_model_state_dict(policy, adapter_weights)
689
+ print(f"[resume] loaded {len(adapter_weights)} tensors from {cfg.resume_adapter}", flush=True)
690
+ except Exception as e:
691
+ print(f"[resume] WARNING: could not load adapter: {e}", flush=True)
692
+ cfg.format_warmup = False
693
+
694
  if cfg.format_warmup:
695
  print("\n=== format warm-start (JSON action behavior) ===", flush=True)
696
  warmup_metrics = run_format_warmup(
 
931
  p.add_argument("--no-save-format-warmup", action="store_true")
932
  p.add_argument("--format-warmup-model-id", default=None)
933
  p.add_argument("--sample-temperature", type=float, default=None)
934
+ p.add_argument("--resume-adapter", default=None)
935
  args = p.parse_args()
936
 
937
  cfg = TrainConfig()
 
965
  cfg.format_warmup_model_id = args.format_warmup_model_id
966
  if args.sample_temperature is not None:
967
  cfg.sample_temperature = args.sample_temperature
968
+ if args.resume_adapter:
969
+ cfg.resume_adapter = args.resume_adapter
970
  if args.no_push:
971
  cfg.push_to_hub = False
972
  if args.no_4bit: