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:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user