From d9a0c7214924ec4d69f364d45f96906c188f30b7 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Tue, 28 Jul 2026 00:22:29 +0800 Subject: [PATCH] feat: store metric logs inside each checkpoint dir, remove log_dir config --- astrai/config/train_config.py | 3 --- astrai/trainer/train_callback.py | 8 ++++---- astrai/trainer/trainer.py | 2 +- scripts/tools/train.py | 8 -------- tests/trainer/conftest.py | 1 - tests/trainer/test_online_e2e.py | 1 - tests/trainer/test_signal_handler.py | 6 ++---- 7 files changed, 7 insertions(+), 22 deletions(-) diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index 12e8de8..dcf9ff6 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -69,9 +69,6 @@ class TrainConfig(BaseConfig): ) # metric setting - log_dir: str = field( - default="./checkpoint/logs", metadata={"help": "Directory for metric logs."} - ) metrics: List[str] = field( default_factory=lambda: ["loss", "lr", "grad_norm"], metadata={"help": "Metrics to record during training."}, diff --git a/astrai/trainer/train_callback.py b/astrai/trainer/train_callback.py index 2bc9866..6cd713f 100644 --- a/astrai/trainer/train_callback.py +++ b/astrai/trainer/train_callback.py @@ -235,7 +235,7 @@ class ProgressBarCallback(TrainCallback): class MetricCallback(TrainCallback): def __init__( self, - log_dir: str, + ckpt_dir: str, save_interval: int, metrics: List[str] = None, val_step: int = 0, @@ -246,8 +246,7 @@ class MetricCallback(TrainCallback): self.val_step = val_step self._next_val_step = 0 - self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs" - self.log_dir.mkdir(parents=True, exist_ok=True) + self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else Path.cwd() / "checkpoint" self.log_cache = [] @@ -306,11 +305,12 @@ class MetricCallback(TrainCallback): @only_on_rank(0) def _flush(self, epoch, step): - log_file = self.log_dir / f"epoch_{epoch}_step_{step}_metric.jsonl" + log_file = self.ckpt_dir / f"epoch_{epoch}_step_{step}" / "metric.jsonl" log_file.parent.mkdir(parents=True, exist_ok=True) with open(log_file, "w") as f: for log in self.log_cache: f.write(json.dumps(log) + "\n") + self.log_cache.clear() def on_optimizer_step(self, context): if ( diff --git a/astrai/trainer/trainer.py b/astrai/trainer/trainer.py index 227da3e..6292081 100644 --- a/astrai/trainer/trainer.py +++ b/astrai/trainer/trainer.py @@ -42,7 +42,7 @@ class Trainer: ), CallbackFactory.create( "metric", - log_dir=cfg.log_dir, + ckpt_dir=cfg.ckpt_dir, save_interval=cfg.ckpt_interval, metrics=cfg.metrics, val_step=cfg.val_step, diff --git a/scripts/tools/train.py b/scripts/tools/train.py index ff765cc..abbb758 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -224,12 +224,6 @@ _START_METHODS = ["spawn", "fork", "forkserver"] default=("loss", "lr", "grad_norm"), help="Metrics to log (repeatable).", ) -@click.option( - "--log_dir", - type=click.Path(), - default="checkpoint/logs", - help="Directory for metric logs.", -) @click.option("--start_epoch", type=int, default=0, help="Start epoch.") @click.option("--start_samples", type=int, default=0, help="Start samples (per rank).") @click.option( @@ -379,7 +373,6 @@ def train( val_split: float, val_step: int, metrics: list[str], - log_dir: str, max_grad_norm: float, random_seed: int, num_workers: int, @@ -538,7 +531,6 @@ def train( val_split=val_split, val_step=val_step, metrics=metrics, - log_dir=log_dir, gradient_checkpointing_modules=grad_ckpt_modules, executor_kwargs=executor_kwargs, extra_kwargs=strategy_kwargs, diff --git a/tests/trainer/conftest.py b/tests/trainer/conftest.py index 87b1cc4..39949c8 100644 --- a/tests/trainer/conftest.py +++ b/tests/trainer/conftest.py @@ -39,7 +39,6 @@ def create_train_config( optimizer_fn=optimizer_fn, scheduler_fn=scheduler_fn, ckpt_dir=test_dir, - log_dir=os.path.join(test_dir, "logs"), n_epoch=n_epoch, batch_per_device=batch_per_device, ckpt_interval=ckpt_interval, diff --git a/tests/trainer/test_online_e2e.py b/tests/trainer/test_online_e2e.py index 7570825..d265e5c 100644 --- a/tests/trainer/test_online_e2e.py +++ b/tests/trainer/test_online_e2e.py @@ -101,7 +101,6 @@ def test_online_dpo_end_to_end(base_test_env): optimizer_fn=optimizer_fn, scheduler_fn=scheduler_fn, ckpt_dir=os.path.join(test_dir, "ckpt"), - log_dir=os.path.join(test_dir, "logs"), n_epoch=1, batch_per_device=2, ckpt_interval=100, diff --git a/tests/trainer/test_signal_handler.py b/tests/trainer/test_signal_handler.py index 9d74f69..d1750d3 100644 --- a/tests/trainer/test_signal_handler.py +++ b/tests/trainer/test_signal_handler.py @@ -50,7 +50,7 @@ class _ReadyCallback: os.fsync(f.fileno()) -def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, log_dir, ready_file): +def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, ready_file): dataset = PicklableDataset() def model_fn(): @@ -71,7 +71,6 @@ def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, log_dir, ready_file): optimizer_fn=optimizer_fn, scheduler_fn=scheduler_fn, ckpt_dir=ckpt_dir, - log_dir=log_dir, n_epoch=1, batch_per_device=batch_per_device, ckpt_interval=ckpt_interval, @@ -86,13 +85,12 @@ def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, log_dir, ready_file): def _spawn_train_and_signal(ckpt_dir, sig, timeout=120): - log_dir = os.path.join(ckpt_dir, "logs") ready_file = os.path.join(ckpt_dir, "ready.txt") ctx = mp.get_context("spawn") p = ctx.Process( target=_inner_run, - args=(2, 1000, ckpt_dir, log_dir, ready_file), + args=(2, 1000, ckpt_dir, ready_file), ) p.start()