feat(paralell): 添加分布式训练配置与并行工具支持

This commit is contained in:
2025-12-05 13:52:17 +08:00
parent d31137a2db
commit d52685facd
4 changed files with 72 additions and 1 deletions
+9
View File
@@ -5,6 +5,7 @@ from torch.utils.data import DataLoader
from khaosz.config import Checkpoint
from khaosz.data import ResumableDistributedSampler
from khaosz.trainer.schedule import BaseScheduler, SchedulerFactory
from khaosz.parallel.utils import get_world_size, get_rank
if TYPE_CHECKING:
from khaosz.trainer.trainer import Trainer
@@ -20,6 +21,9 @@ class TrainContext:
batch_iter: int = field(default=0)
loss: float = field(default=0.0)
wolrd_size: int = field(default=1)
rank: int = field(default=0)
def asdict(self) -> dict:
return {field.name: getattr(self, field.name)
for field in fields(self)}
@@ -102,4 +106,9 @@ class TrainContextBuilder:
return self
def build(self) -> TrainContext:
if self.trainer.train_config.nprocs > 1:
self._context.wolrd_size = get_world_size()
self._context.rank = get_rank()
return self._context