refactor: merge validation into MetricCallback, simplify progress bar to optimizer steps

- Remove separate ValidationCallback, merge into MetricCallback
- Progress bar now tracks optimizer steps instead of micro-steps
- Remove unused log_interval config field and CLI flag
- Fix validation all_reduce: use SUM(loss, count) instead of AVG
- Simplify metric logging: always log every optimizer step
- Add grad_norm display to progress bar
This commit is contained in:
2026-07-03 21:43:08 +08:00
parent dfb151537b
commit 70c0e5de90
4 changed files with 56 additions and 78 deletions
-8
View File
@@ -159,12 +159,6 @@ def parse_args() -> argparse.Namespace:
default="checkpoint/logs",
help="Directory for metric logs.",
)
parser.add_argument(
"--log_interval",
type=int,
default=1,
help="Number of optimizer steps between metric logs.",
)
parser.add_argument(
"--grpo_sync_interval",
type=int,
@@ -329,7 +323,6 @@ def train(
val_step: int,
metrics: list[str],
log_dir: str,
log_interval: int,
dpo_beta: float,
grpo_clip_eps: float,
grpo_kl_coef: float,
@@ -465,7 +458,6 @@ def train(
val_step=val_step,
metrics=metrics,
log_dir=log_dir,
log_interval=log_interval,
gradient_checkpointing_modules=grad_ckpt_modules,
executor_kwargs=executor_kwargs,
extra_kwargs=strategy_kwargs,