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
+2 -10
View File
@@ -8,7 +8,6 @@ from torch.optim import AdamW
from khaosz.config import ModelParameter, TrainConfig, CosineScheduleConfig
from khaosz.trainer import Trainer, SchedulerFactory
from khaosz.data import DatasetLoader
from khaosz.parallel import get_current_device, spawn_parallel_fn
def parse_args() -> argparse.Namespace:
@@ -96,8 +95,6 @@ def train(
window_size = parameter.config.m_len
model = parameter.model
current_device = get_current_device()
model = fsdp_wrap(model.to(device=current_device, dtype=torch.bfloat16))
kwargs = {
"dpo_beta": dpo_beta,
@@ -150,6 +147,7 @@ def train(
pin_memory=pin_memory,
nprocs=nprocs,
extra_kwargs=kwargs,
parallel_fn=fsdp_wrap
)
trainer = Trainer(train_config)
@@ -158,10 +156,4 @@ def train(
if __name__ == "__main__":
args = parse_args()
spawn_parallel_fn(
train,
world_size=args.nprocs,
backend="nccl",
**vars(args)
)
train(**vars(args))