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
+17 -1
View File
@@ -4,7 +4,7 @@ from torch.optim import Optimizer
from torch.optim.lr_scheduler import LRScheduler
from dataclasses import dataclass, field
from typing import Optional
from typing import Callable, Optional
@dataclass
@@ -88,6 +88,22 @@ class TrainConfig:
default=1,
metadata={"help": "Number of processes for distributed training."}
)
backend: str = field(
default="nccl",
metadata={"help": "Distributed training backend."}
)
master_addr: str = field(
default="localhost",
metadata={"help": "Master address for distributed training."}
)
master_port: str = field(
default="29500",
metadata={"help": "Master port for distributed training."}
)
parallel_fn: Optional[Callable] = field(
default=None,
metadata={"help": "Parallel function for training."}
)
# others
extra_kwargs: dict = field(