fix(trainer): 修复参数传递问题和检查点保存问题
This commit is contained in:
@@ -93,4 +93,13 @@ class TrainConfig:
|
||||
kwargs: dict = field(
|
||||
default_factory=dict,
|
||||
metadata={"help": "Other arguments."}
|
||||
)
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
self.validate()
|
||||
|
||||
def validate(self):
|
||||
required_fields = ["model", "strategy", "dataset", "optimizer", "scheduler"]
|
||||
for field_name in required_fields:
|
||||
if getattr(self, field_name) is None:
|
||||
raise ValueError(f"{field_name} is required.")
|
||||
|
||||
@@ -94,18 +94,17 @@ class CheckpointCallback(TrainCallback):
|
||||
def __init__(self, interval: int, save_dir: str):
|
||||
self.interval = interval
|
||||
self.save_dir = save_dir
|
||||
self.checkpoint = None
|
||||
self.last_ckpt_iter = 0
|
||||
|
||||
def _save_checkpoint(self, context: 'TrainContext'):
|
||||
save_path = os.path.join(self.save_dir, f"epoch_{context.epoch}iter_{context.iteration}")
|
||||
self.checkpoint = Checkpoint(
|
||||
context.checkpoint = Checkpoint(
|
||||
context.optimizer.state_dict(),
|
||||
context.scheduler.state_dict(),
|
||||
context.epoch,
|
||||
context.iteration
|
||||
)
|
||||
self.checkpoint.save(save_path)
|
||||
context.checkpoint.save(save_path)
|
||||
self.last_ckpt_iter = context.iteration
|
||||
|
||||
def on_batch_end(self, context: 'TrainContext'):
|
||||
|
||||
Reference in New Issue
Block a user