feat: store metric logs inside each checkpoint dir, remove log_dir config
This commit is contained in:
@@ -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."},
|
||||||
|
|||||||
@@ -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 (
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user