feat: store metric logs inside each checkpoint dir, remove log_dir config

This commit is contained in:
2026-07-28 00:22:29 +08:00
parent 5ab18bec48
commit d9a0c72149
7 changed files with 7 additions and 22 deletions
-3
View File
@@ -69,9 +69,6 @@ class TrainConfig(BaseConfig):
) )
# metric setting # metric setting
log_dir: str = field(
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
)
metrics: List[str] = field( metrics: List[str] = field(
default_factory=lambda: ["loss", "lr", "grad_norm"], default_factory=lambda: ["loss", "lr", "grad_norm"],
metadata={"help": "Metrics to record during training."}, metadata={"help": "Metrics to record during training."},
+4 -4
View File
@@ -235,7 +235,7 @@ class ProgressBarCallback(TrainCallback):
class MetricCallback(TrainCallback): class MetricCallback(TrainCallback):
def __init__( def __init__(
self, self,
log_dir: str, ckpt_dir: str,
save_interval: int, save_interval: int,
metrics: List[str] = None, metrics: List[str] = None,
val_step: int = 0, val_step: int = 0,
@@ -246,8 +246,7 @@ class MetricCallback(TrainCallback):
self.val_step = val_step self.val_step = val_step
self._next_val_step = 0 self._next_val_step = 0
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs" self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else Path.cwd() / "checkpoint"
self.log_dir.mkdir(parents=True, exist_ok=True)
self.log_cache = [] self.log_cache = []
@@ -306,11 +305,12 @@ class MetricCallback(TrainCallback):
@only_on_rank(0) @only_on_rank(0)
def _flush(self, epoch, step): 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) log_file.parent.mkdir(parents=True, exist_ok=True)
with open(log_file, "w") as f: with open(log_file, "w") as f:
for log in self.log_cache: for log in self.log_cache:
f.write(json.dumps(log) + "\n") f.write(json.dumps(log) + "\n")
self.log_cache.clear()
def on_optimizer_step(self, context): def on_optimizer_step(self, context):
if ( if (
+1 -1
View File
@@ -42,7 +42,7 @@ class Trainer:
), ),
CallbackFactory.create( CallbackFactory.create(
"metric", "metric",
log_dir=cfg.log_dir, ckpt_dir=cfg.ckpt_dir,
save_interval=cfg.ckpt_interval, save_interval=cfg.ckpt_interval,
metrics=cfg.metrics, metrics=cfg.metrics,
val_step=cfg.val_step, val_step=cfg.val_step,
-8
View File
@@ -224,12 +224,6 @@ _START_METHODS = ["spawn", "fork", "forkserver"]
default=("loss", "lr", "grad_norm"), default=("loss", "lr", "grad_norm"),
help="Metrics to log (repeatable).", 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_epoch", type=int, default=0, help="Start epoch.")
@click.option("--start_samples", type=int, default=0, help="Start samples (per rank).") @click.option("--start_samples", type=int, default=0, help="Start samples (per rank).")
@click.option( @click.option(
@@ -379,7 +373,6 @@ def train(
val_split: float, val_split: float,
val_step: int, val_step: int,
metrics: list[str], metrics: list[str],
log_dir: str,
max_grad_norm: float, max_grad_norm: float,
random_seed: int, random_seed: int,
num_workers: int, num_workers: int,
@@ -538,7 +531,6 @@ def train(
val_split=val_split, val_split=val_split,
val_step=val_step, val_step=val_step,
metrics=metrics, metrics=metrics,
log_dir=log_dir,
gradient_checkpointing_modules=grad_ckpt_modules, gradient_checkpointing_modules=grad_ckpt_modules,
executor_kwargs=executor_kwargs, executor_kwargs=executor_kwargs,
extra_kwargs=strategy_kwargs, extra_kwargs=strategy_kwargs,
-1
View File
@@ -39,7 +39,6 @@ def create_train_config(
optimizer_fn=optimizer_fn, optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn, scheduler_fn=scheduler_fn,
ckpt_dir=test_dir, ckpt_dir=test_dir,
log_dir=os.path.join(test_dir, "logs"),
n_epoch=n_epoch, n_epoch=n_epoch,
batch_per_device=batch_per_device, batch_per_device=batch_per_device,
ckpt_interval=ckpt_interval, ckpt_interval=ckpt_interval,
-1
View File
@@ -101,7 +101,6 @@ def test_online_dpo_end_to_end(base_test_env):
optimizer_fn=optimizer_fn, optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn, scheduler_fn=scheduler_fn,
ckpt_dir=os.path.join(test_dir, "ckpt"), ckpt_dir=os.path.join(test_dir, "ckpt"),
log_dir=os.path.join(test_dir, "logs"),
n_epoch=1, n_epoch=1,
batch_per_device=2, batch_per_device=2,
ckpt_interval=100, ckpt_interval=100,
+2 -4
View File
@@ -50,7 +50,7 @@ class _ReadyCallback:
os.fsync(f.fileno()) 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() dataset = PicklableDataset()
def model_fn(): 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, optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn, scheduler_fn=scheduler_fn,
ckpt_dir=ckpt_dir, ckpt_dir=ckpt_dir,
log_dir=log_dir,
n_epoch=1, n_epoch=1,
batch_per_device=batch_per_device, batch_per_device=batch_per_device,
ckpt_interval=ckpt_interval, 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): 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") ready_file = os.path.join(ckpt_dir, "ready.txt")
ctx = mp.get_context("spawn") ctx = mp.get_context("spawn")
p = ctx.Process( p = ctx.Process(
target=_inner_run, target=_inner_run,
args=(2, 1000, ckpt_dir, log_dir, ready_file), args=(2, 1000, ckpt_dir, ready_file),
) )
p.start() p.start()