refactor(trainer): 统一参数命名以提升可读性

This commit is contained in:
2025-09-28 22:14:24 +08:00
parent fa43ed2943
commit 1c9063fd3d
4 changed files with 37 additions and 37 deletions
+4 -4
View File
@@ -190,7 +190,7 @@ class TrainConfig:
default=None,
metadata={"help": "Optimizer for training."}
)
ckpt_dir: str = field(
checkpoint_dir: str = field(
default="./checkpoint",
metadata={"help": "Checkpoint directory."}
)
@@ -202,11 +202,11 @@ class TrainConfig:
default=4,
metadata={"help": "Batch size for training."}
)
n_iter_ckpt: int = field(
checkpoint_interval: int = field(
default=5000,
metadata={"help": "Number of iterations between checkpoints."}
)
n_iter_step: int = field(
accumulation_steps: int = field(
default=1,
metadata={"help": "Number of iterations between steps."}
)
@@ -256,7 +256,7 @@ class ScheduleConfig(ABC):
@dataclass
class CosineScheduleConfig(ScheduleConfig):
total_steps: int = field( # 更准确的命名
total_steps: int = field(
default=None,
metadata={"help": "Total training steps for cosine schedule."}
)
+3 -3
View File
@@ -32,7 +32,7 @@ class Trainer:
train_config: TrainConfig
):
current_iter = len(loss_list)
save_path = os.path.join(train_config.ckpt_dir, f"iter_{current_iter}")
save_path = os.path.join(train_config.checkpoint_dir, f"iter_{current_iter}")
self.checkpoint.loss_list = loss_list
self.checkpoint.optim_state = train_config.optimizer.state_dict()
self.checkpoint.save(save_path)
@@ -93,7 +93,7 @@ class Trainer:
#backward
loss.backward()
#step
if current_iter % train_config.n_iter_step == 0:
if current_iter % train_config.accumulation_steps == 0:
clip_grad_norm_(
self.checkpoint.model.parameters(),
train_config.max_grad_norm
@@ -108,7 +108,7 @@ class Trainer:
"lr": f"{train_config.optimizer.param_groups[0]['lr']:.2e}"
})
#save checkpotint
if current_iter - last_ckpt_iter >= train_config.n_iter_ckpt:
if current_iter - last_ckpt_iter >= train_config.checkpoint_interval:
self.save_checkpoint(loss_list, train_config)
last_ckpt_iter = current_iter