已开启
【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 9 local_map 扩展(可选) #147
changzherui创建于  5月15日
changzherui
changzherui成员
5月15日 创建

Part 9 — local_map 扩展(可选)

仅当首个需要 q/k/v 走 head-shard 路径的模型(如某些 attention 变种、GatedDeltaNet 的 TP 切分等)迁移时才做。与主线解耦,可独立推进,也可推迟。


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 接口差异说明

# 差异点 原因
1 local_map.py 是 hyper 自己的包装;torchtitan 直接用 torch.distributed.tensor.experimental.local_map hyper 跨后端,MindSpore 必须先 raise
2 MindSpore 后端可能永久 raise(除非未来 MindSpore 提供等价 API) MindSpore 当前 DTensor / Layout 体系没有等价的"local 张量解包" API
3 LocalMapConfig.callable_path 用字符串标识 inner callable(如 "inner_sdpa"),M9 通过 _unwrap_inner_callable 解析;torchtitan 直接传 callable tyro / 序列化友好;callable 不能直接 dataclass field

5. 开发步骤

步 内容 工期
1 core/dtensor/local_map.py 跨后端包装 + UT 0.5 d
2 改 protocols/module.py:把占位 NotImplementedError 替换为真实调用;新增 _unwrap_inner_callable / _bind_inner_callable 1 d
3 改 models/common/decoder_sharding.py:新增 set_gqa_inner_attention_local_map 0.5 d
4 写 2-card 玩具 attention 数值对齐测试 1 d

6. 验证标准

新建 tests/torch/ut/dtensor/test_local_map.py 和 tests/torch/st/local_map_attention/:

测试 断言要点
test_local_map.py 2-card:把 f(x: DTensor[Shard(0)], y: DTensor[Replicate()]) -> DTensor[Shard(0)] 包成 local_map;调用后输入展开为 local 张量、输出还原 DTensor;数值与单卡 f(x_full, y_full) 一致
test_local_map_grad.py 反向梯度 Shard(0) 正确
test_local_map_mindspore_raises.py MindSpore 后端调 local_map 必须 raise NotImplementedError("...pending...")
tests/torch/st/local_map_attention/test_sdpa_tp.py 8-card:q/k/v Shard(2) → SDPA local → Shard(2);与单卡数值一致(误差 ≤ 1e-5)

通过门槛:

  • 4 个测试全绿。
  • M9 启用后,已有 M6 / M8 模型行为不变(local_map 字段默认 None,路径无差异)。

7. 工期 & 依赖

工期 3 天
依赖 M1
触发条件 仅当首个需要 q/k/v 走 head-shard 的模型迁移时执行;否则可永久推迟
下游 视模型需求
likedislike
changzheruichangzherui成员
5月15日 修改标题为 “hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M9”
changzheruichangzherui成员
5月15日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”,原标题为“hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”
changzheruichangzherui成员
5月25日 修改标题为 “【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— Part 9 local_map 扩展(可选)”,原标题为“【RFC】hyper-parallel 支持torchtitan 风格 Module/Config —— M9 local_map 扩展(可选)”
changzheruichangzherui成员
5月25日 修改了issue 的描述