feat: GRPO CLI 接入 + on-policy,OpenAI API top_k 参数化,补充训练参数表

- train.py 新增 --train_type=grpo 及参数 (--grpo_clip_eps, --grpo_kl_coef, --group_size, --grpo_sync_interval, --start_epoch)
- GRPOStrategy 统一 on-policy 模式,ratio = exp(logπ_θ - logπ_ref),PPO 裁剪目标,sync_interval 自动同步 ref_model
- ChatCompletionRequest 新增 top_k 参数,不再硬编码
- 补充 README 完整训练参数表(含此前缺失的 max_grad_norm / adamw / window_size / stride 等)
This commit is contained in:
2026-05-09 12:22:33 +08:00
parent 34a511e36e
commit bc7c82977e
5 changed files with 91 additions and 19 deletions
+33 -6
View File
@@ -23,7 +23,7 @@ def parse_args() -> argparse.Namespace:
"--train_type",
type=str,
required=True,
choices=["seq", "sft", "dpo"],
choices=["seq", "sft", "dpo", "grpo"],
help="Train type.",
)
parser.add_argument(
@@ -42,9 +42,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument(
"--n_epoch", type=int, default=1, help="Number of epochs to train."
)
parser.add_argument(
"--batch_size", type=int, default=1, help="Batch size for training."
)
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
parser.add_argument(
"--accumulation_steps",
type=int,
@@ -106,6 +104,17 @@ def parse_args() -> argparse.Namespace:
"--stride", type=int, default=None, help="the step size of the input sequence."
)
parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.")
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
parser.add_argument(
"--on_policy",
action="store_true",
default=False,
help="Enable on-policy GRPO mode.",
)
parser.add_argument(
"--grpo_kl_coef", type=float, default=0.01, help="GRPO KL penalty coefficient."
)
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
parser.add_argument(
"--label_smoothing",
type=float,
@@ -125,6 +134,13 @@ def parse_args() -> argparse.Namespace:
default="checkpoint",
help="Directory to save checkpoints.",
)
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
parser.add_argument(
"--grpo_sync_interval",
type=int,
default=200,
help="GRPO ref model sync interval (steps).",
)
parser.add_argument(
"--start_epoch", type=int, default=0, help="Start epoch for training."
)
@@ -182,6 +198,10 @@ def train(
ckpt_interval: int,
ckpt_dir: str,
dpo_beta: float,
grpo_clip_eps: float,
grpo_kl_coef: float,
group_size: int,
grpo_sync_interval: int,
adamw_beta1: float,
adamw_beta2: float,
adamw_weight_decay: float,
@@ -195,7 +215,7 @@ def train(
nprocs: int,
device_type: str,
):
assert train_type in ["seq", "sft", "dpo"]
assert train_type in ["seq", "sft", "dpo", "grpo"]
assert os.path.exists(param_path)
# Load config
@@ -216,7 +236,14 @@ def train(
state_dict = st.load_file(weights_path)
model.load_state_dict(state_dict, strict=False)
strategy_kwargs = {"dpo_beta": dpo_beta, "label_smoothing": label_smoothing}
strategy_kwargs = {
"dpo_beta": dpo_beta,
"label_smoothing": label_smoothing,
"clip_eps": grpo_clip_eps,
"kl_coef": grpo_kl_coef,
"group_size": group_size,
"sync_interval": grpo_sync_interval,
}
dataset = DatasetFactory.load(
train_type=train_type,