refactor(trainer): 统一参数命名以提升可读性
This commit is contained in:
@@ -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."}
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user