已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 1 协议层 #139
changzherui创建于 5月15日
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M1”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— step1”
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M1”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— step1”
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M1 协议层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M1”
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M1 协议层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M1”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M1 协议层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M1 协议层”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M1 协议层”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M1 协议层”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 1 协议层”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M1 协议层”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 1 协议层”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M1 协议层”
5月25日 修改了issue 的描述
Part 1 —
hyper_parallel/protocols/协议层1. 目标
实现 torchtitan 的
Configurable / Module / BaseModel / ShardingConfig / MeshAxisName / ModelSpec一整套协议,作为后续 M2–M9 共同的基础抽象层。2. 任务边界(新增文件)
新增包目录
hyper_parallel/protocols/:types.pyMeshAxisName(StrEnum):定义DP / DP_REPLICATE / DP_SHARD / FSDP / TP / CP / PP / EP / EFSDP。字面量与 hyper 现用 mesh 维度名对齐("fsdp" / "dp_shard" / "tp" / "cp" / "pp" / "ep" / "dp_replicate")torchtitan/protocols/types.pyconfigurable.pyConfigurable基类 + 嵌套Configurable.Config(@dataclass(kw_only=True, slots=True));__init_subclass__自动把Config._owner = cls,Config.build(**kwargs)即构造cls(self, **kwargs);提供replace / traverse / to_dicttorchtitan/config/configurable.pysharding.pyNamedPlacement = dict[MeshAxisName, Placement];@dataclass class ShardingConfig(state_shardings / in_src_shardings / in_dst_shardings / out_dst_shardings / local_map);LocalMapConfig;resolve_placements(named, mesh_axis_names) -> list[Placement]。Placement直接复用hyper_parallel.core.dtensor.placement_types.{Placement, Shard, Replicate, Partial}torchtitan/protocols/sharding.pymodule.pyclass Module(platform.Module, Configurable):实现init_states / _init_self_parameters / _init_self_buffers / _init_param / _cache_pos_arg_names / parallelize / _shard_inputs / _shard_outputs / from_nn_module;定义ModuleList / ModuleDict / Sequential容器版torchtitan/protocols/module.pymodel.pyclass BaseModel(Module)+BaseModel.Config(Module.Config),含verify_module_protocol / init_weights / update_from_config(trainer_config) / get_nparams_and_flops;ModelConfigConverter(Configurable)占位torchtitan/protocols/model.pystate_dict_adapter.pyclass BaseStateDictAdapter(ABC),方法签名兼容现有hyper_parallel/models/spec/state_dict_adapter.py:28的Protocol(load_hf_state_dict / save_hf_state_dict)torchtitan/protocols/state_dict_adapter.pymodel_spec.pyModelSpec:兼容现有models/spec/model_spec.py:23的 5 字段,新增model: BaseModel.Config | None;旧字段build_model_fn改为Optional,注册时二选一torchtitan/protocols/model_spec.py__init__.py3. 核心设计点
3.1
Configurable.__init_subclass__class Configurable: Config: type["Configurable.Config"] @dataclass(kw_only=True, slots=True) class Config: _owner: ClassVar[type["Configurable"]] = None def build(self, **runtime_kwargs) -> "Configurable": return self._owner(self, **runtime_kwargs) def __init_subclass__(cls, **kw): super().__init_subclass__(**kw) if "Config" in cls.__dict__: cls.Config._owner = cls3.2
Module.parallelize(tp_mesh)流程def parallelize(self, tp_mesh): sc = self.sharding_config if sc is None: # 递归 children for child in self.children(): if isinstance(child, Module): child.parallelize(tp_mesh) return if sc.local_map is not None: raise NotImplementedError("local_map will be added in M9") # 1. 分布参数 / buffer for path, named_p in sc.state_shardings.items(): param = _get_attr_by_path(self, path) placements = resolve_placements(named_p, tp_mesh.mesh_dim_names) new_local = distribute_tensor(param.data, tp_mesh, placements) _set_param_by_path(self, path, platform.Parameter(new_local)) # 2. 包 forward 做 in/out reshard self._wrap_forward_with_reshard(tp_mesh) # 3. 递归子模块 for child in self.children(): if isinstance(child, Module): child.parallelize(tp_mesh)注意
distribute_tensor来自hyper_parallel.core.dtensor.dtensor:409,签名(tensor, device_mesh, placements),与 torch 版兼容。3.3
resolve_placementsdef resolve_placements( named: NamedPlacement, # dict[MeshAxisName, Placement] mesh_axis_names: tuple[str, ...], # 实际 mesh 的维度名 ) -> list[Placement]: out = [] for axis_name in mesh_axis_names: if axis_name in named: out.append(named[axis_name]) else: out.append(Replicate()) # 未声明的轴默认 Replicate # 缺主轴报错 declared = set(named.keys()) extra = declared - set(mesh_axis_names) if extra: raise ValueError(f"NamedPlacement has axes not in mesh: {extra}") return out4. 与 torchtitan 接口差异说明
Module继承(platform.Module, Configurable)而非(nn.Module, Configurable)distribute_tensor用hyper_parallel.core.dtensor.distribute_tensor(core/dtensor/dtensor.py:409)MeshAxisName字面量优先 hyper 现用名(FSDP="fsdp"),同时收DP_SHARD="dp_shard"trainer/parallel_dims.py已建的 mesh 命名对齐ShardingConfig.local_map字段保留但 raise NotImplementedErrorConfigurable.__init_subclass__的slots=True校验只对Module子树生效slots,强制开启会破坏存量代码ModelSpecv2 兼容 v1:build_model_fn / model: BaseModel.Config二选一5. 开发步骤
types.py+__init__.py。configurable.py。__init_subclass__自动绑 owner。Config.build()调cls(self, **runtime_kwargs)。Configurable.Config.replace(**kw)用dataclasses.replace实现。sharding.py。resolve_placements是核心。module.py的Module基类 +ModuleList / ModuleDict / Sequential。model.py的BaseModel。update_from_config默认空实现,由子类重写。state_dict_adapter.py+model_spec.py。6. 验证标准
新建
tests/torch/ut/protocols/:test_configurable.pyConfig.build()构造正确;replace(field=val)不改原对象;traverse遍历嵌套Config子树;to_dict输出可 yaml 序列化test_module.pyModule子类的init_states()、_init_self_buffers(device)、from_nn_module(nn.Linear)复用同一份Config;_cache_pos_arg_names通过inspect.signature(self.forward)缓存test_module_parallelize.pynn.Linear风格的Module,配state_shardings={"weight": {"tp": Shard(0)}}、in_dst_shardings={"input": {"tp": Replicate()}}、out_dst_shardings={"output": {"tp": Shard(-1)}};module.parallelize(mesh)后参数变 DTensor、forward 通过、数值与单卡一致test_sharding_resolve.pyresolve_placements对缺轴默认Replicate();对额外轴 raiseValueError;按mesh.mesh_dim_names顺序输出test_module_local_map_stub.pylocal_map非空时调parallelize必须 raiseNotImplementedError("local_map will be added in M9")通过门槛:
import hyper_parallel.protocols不引入循环依赖。tests/torch/st/qwen3_5/等存量测试 0 受影响。hyper_parallel/__init__.py不暴露新符号。7. 工期 & 依赖