v6: GRPO with --resume-adapter to start from SFT checkpoint
Browse files- 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 |
-
# -----
|
| 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:
|