[v1] refactor registry plugin structure and params (#10641)

This commit is contained in:
Jiaqi
2026-07-24 15:23:21 +08:00
committed by GitHub
parent 19e9fe3ced
commit 3f77101580
50 changed files with 843 additions and 1007 deletions

View File

@@ -29,16 +29,20 @@ And data parallelism types:
from dataclasses import dataclass
from datetime import timedelta
from enum import StrEnum
from typing import Any, Optional
from typing import TYPE_CHECKING, Any, Optional
from torch.distributed import barrier, destroy_process_group, init_process_group
from torch.distributed.device_mesh import DeviceMesh, init_device_mesh
from ..utils import logging
from ..utils.types import DistributedConfig, ProcessGroup, TensorLike
from ..utils.types import ProcessGroup, TensorLike
from . import helper
if TYPE_CHECKING:
from ..config.training_args import TrainingArguments
logger = logging.get_logger(__name__)
@@ -128,12 +132,13 @@ class DistributedInterface:
return cls._instance
def __init__(self, config: DistributedConfig | None = None) -> None:
def __init__(
self,
training_args: "TrainingArguments | None" = None,
) -> None:
if self._initialized:
return
self.dist_config = config
helper.set_device_index()
self._is_distributed = helper.is_distributed()
self._rank = helper.get_rank()
@@ -143,17 +148,17 @@ class DistributedInterface:
self.current_device = helper.get_current_device()
self.device_count = helper.get_device_count()
if config is None:
if training_args is None:
self.strategy = DistributedStrategy()
timeout = 18000
else:
self.strategy = DistributedStrategy(
mp_replicate_size=config.get("mp_replicate_size", 1),
mp_shard_size=config.get("mp_shard_size", None),
dp_size=config.get("dp_size", None),
cp_size=config.get("cp_size", 1),
mp_replicate_size=training_args.mp_replicate_size,
mp_shard_size=training_args.mp_shard_size,
dp_size=training_args.dp_size,
cp_size=training_args.cp_size,
)
timeout = config.get("timeout", 18000)
timeout = training_args.dist_timeout
if self._is_distributed:
init_process_group(timeout=timedelta(seconds=timeout), backend=helper.get_process_group_backend())