feat: metric 参数通过 TrainConfig 传递

- TrainConfig 新增 log_dir/log_interval/metrics 配置字段

- metric_logger 调用改用 **kwargs 传递,BaseFactory.create 自动过滤
This commit is contained in:
2026-05-19 17:50:24 +08:00
parent e0a3337c22
commit 45479b5731
2 changed files with 21 additions and 2 deletions
+14 -1
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field, fields
from typing import Callable, Optional
from typing import Callable, List, Optional
import torch.nn as nn
from torch.optim import Optimizer
@@ -56,6 +56,19 @@ class TrainConfig(BaseConfig):
default=5000, metadata={"help": "Number of iterations between checkpoints."}
)
# metric setting
log_dir: str = field(
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
)
log_interval: int = field(
default=100,
metadata={"help": "Number of batch iterations between metric logs."},
)
metrics: List[str] = field(
default_factory=lambda: ["loss", "lr"],
metadata={"help": "Metrics to record during training."},
)
# dataloader setting
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
num_workers: int = field(