feat(trainer): 支持分布式训练配置与检查点加载优化

This commit is contained in:
2025-12-19 19:34:39 +08:00
parent eab7a51bb6
commit 573f041c51
8 changed files with 67 additions and 27 deletions
+7 -2
View File
@@ -1,3 +1,4 @@
import os
import torch
import numpy as np
from khaosz.config import *
@@ -31,10 +32,14 @@ def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
checkpoint = None
try:
checkpoint = trainer.train()
assert checkpoint.iteration == 2
except Exception:
# Handle any exceptions
pass
checkpoint = trainer.train(checkpoint)
load_dir = os.path.join(base_test_env["test_dir"], "epoch_0_iter_2")
checkpoint = Checkpoint.load(load_dir)
trainer.train(checkpoint)
load_dir = os.path.join(base_test_env["test_dir"], "epoch_1_iter_10")
checkpoint = Checkpoint.load(load_dir)
assert checkpoint.iteration == 10