已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 9 local_map 扩展(可选) #147
changzherui创建于 5月15日
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M9”
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M9”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 9 local_map 扩展(可选)”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 9 local_map 扩展(可选)”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”
5月25日 修改了issue 的描述
Part 9 —
local_map扩展(可选)1. 目标
为
Module.parallelize增加 torchtitan 风格的local_map支持 —— 把 sharded DTensor 输入解包为 local 张量、调用纯 local 函数(如 SDPA / attention kernel)、再包回 DTensor。2. 任务边界
hyper_parallel/core/dtensor/local_map.py(新)torch.distributed.tensor.experimental.local_map:torch 后端直接转发;mindspore 后端raise NotImplementedError("local_map mindspore support pending")hyper_parallel/protocols/module.py(改)Module.parallelize中sharding_config.local_map is not None时启用 —— 把 M1 时的占位NotImplementedError替换为真实调用hyper_parallel/models/common/decoder_sharding.py(改)set_gqa_inner_attention_local_map(inner_attn_cfg, *, return_lse)助手,给q/k/v Shard(2)(head 维度)→ 进 SDPA → 出Shard(2)3. 核心设计
3.1 跨后端
local_map包装# hyper_parallel/core/dtensor/local_map.py from hyper_parallel.platform import get_platform from hyper_parallel.platform.platform import PlatformType platform = get_platform() def local_map( func, out_placements, in_placements, in_grad_placements=None, device_mesh=None, *, redistribute_inputs=True, ): """Cross-backend wrapper around torch's experimental local_map.""" if platform.platform_type == PlatformType.PYTORCH: # pylint: disable=C0415 from torch.distributed.tensor.experimental import local_map as torch_local_map return torch_local_map( func, out_placements=out_placements, in_placements=in_placements, in_grad_placements=in_grad_placements, device_mesh=device_mesh, redistribute_inputs=redistribute_inputs, ) if platform.platform_type == PlatformType.MINDSPORE: raise NotImplementedError( "local_map on MindSpore backend is pending. " "Currently the inner-attention TP path requires PyTorch backend." ) raise RuntimeError(f"Unknown platform: {platform.platform_type}")3.2
Module.parallelize启用local_map把 M1 中的占位逻辑:
if sc.local_map is not None: raise NotImplementedError("local_map will be added in M9")替换为:
if sc.local_map is not None: from hyper_parallel.core.dtensor.local_map import local_map inner_fn = self._unwrap_inner_callable(sc.local_map.callable_path) wrapped = local_map( inner_fn, out_placements=[ resolve_placements(p, tp_mesh.mesh_dim_names) for p in sc.local_map.out_placements ], in_placements=[ resolve_placements(p, tp_mesh.mesh_dim_names) for p in sc.local_map.in_placements ], device_mesh=tp_mesh, redistribute_inputs=True, ) self._bind_inner_callable(sc.local_map.callable_path, wrapped)3.3
set_gqa_inner_attention_local_map助手def set_gqa_inner_attention_local_map( inner_attn_cfg, *, return_lse: bool = False, ): """Mark inner attention kernel (e.g. F.scaled_dot_product_attention) as a local-map region: q/k/v come in as Shard(2) on head dim, SDPA runs on local heads, output comes back as Shard(2). """ inner_attn_cfg.sharding_config = ShardingConfig( local_map=LocalMapConfig( callable_path="inner_sdpa", in_placements=( {MeshAxisName.TP: Shard(2)}, # q {MeshAxisName.TP: Shard(2)}, # k {MeshAxisName.TP: Shard(2)}, # v ), out_placements=( ({MeshAxisName.TP: Shard(2)},) # attn_output if not return_lse else ( {MeshAxisName.TP: Shard(2)}, {MeshAxisName.TP: Shard(2)}, # lse ) ), ) )4. 与 torchtitan 接口差异说明
local_map.py是 hyper 自己的包装;torchtitan 直接用torch.distributed.tensor.experimental.local_mapLocalMapConfig.callable_path用字符串标识 inner callable(如"inner_sdpa"),M9 通过_unwrap_inner_callable解析;torchtitan 直接传 callable5. 开发步骤
core/dtensor/local_map.py跨后端包装 + UTprotocols/module.py:把占位 NotImplementedError 替换为真实调用;新增_unwrap_inner_callable / _bind_inner_callablemodels/common/decoder_sharding.py:新增set_gqa_inner_attention_local_map6. 验证标准
新建
tests/torch/ut/dtensor/test_local_map.py和tests/torch/st/local_map_attention/:test_local_map.pyf(x: DTensor[Shard(0)], y: DTensor[Replicate()]) -> DTensor[Shard(0)]包成 local_map;调用后输入展开为 local 张量、输出还原 DTensor;数值与单卡f(x_full, y_full)一致test_local_map_grad.pyShard(0)正确test_local_map_mindspore_raises.pylocal_map必须 raiseNotImplementedError("...pending...")tests/torch/st/local_map_attention/test_sdpa_tp.pyq/k/v Shard(2)→ SDPA local →Shard(2);与单卡数值一致(误差 ≤ 1e-5)通过门槛:
7. 工期 & 依赖
q/k/v走 head-shard 的模型迁移时执行;否则可永久推迟