已开启
[RFC]: pytorch for ascend 支持symmetric memory #2446
zzzzzzk创建于  6月22日
zzzzzzk
6月22日 创建

状态(Status): Draft
作者(Authors): zhangzekun11@huawei.com
创建日期(Created): 20260622
**更新日期(Updated):**20260622
相关 Issue/PR: #123(关联 Issue/PR 以便追踪背景)

1. 概述

1.1 简介

torch_npu Symmetric Memory(对称内存)是一套内存抽象与通信机制,其设计理念源自 SHMEM/PGAS 编程模型。所谓"对称",是指所有参与通信的 rank 在完全相同的虚拟地址上分配大小相同的缓冲区——对任意一段数据,rank i 的地址 A 与 rank j 的地址 A 指向逻辑上相同的通信槽位。这一特性通过 CANN 虚拟内存管理 API 实现:分配时建立物理内存与虚拟地址的映射,注册时 HCCL 将各 rank 的物理内存映射到统一的对称虚拟地址空间,从而实现跨进程的内存共享。

  • 为每个 rank 申请虚拟内存,假设每个 rank 对应的虚拟内存大小为 heap_size,通信域中的 rank 数量为 rank_size,则通信域中总虚拟内存大小为 heap_size * rank_size。
  • 每个 rank 的虚拟地址布局都是相同的。
  • 不同 rank 的物理地址映射到每个 rank 对应位置的虚拟地址,实现对其他 rank 内存的访问。

对称内存功能使得 HCCL 可以直接对业务传入的内存进行操作,无需经过中间缓冲区(HCCL buffer),从而减少内存拷贝开销。
image.png

1.2 动机

  1. 传统集合通信路径开销大。 传统 HCCL 内核中,各 rank 缓冲区地址互不相同,通信前需要交换地址、建立映射,内核内部还要经过多级同步。对于小消息,这些固定开销远超数据传输本身,导致延迟居高不下、带宽利用率低下。
  2. 推理场景对小消息延迟极度敏感。 大模型推理(尤其是 MoE 模型)的通信负载以小消息为主:专家路由产生的 token 分发与合并、张量并行中的频繁 AllReduce 等。这类场景的瓶颈是延迟而非带宽,传统内核的固定同步开销被进一步放大,成为端到端推理延迟的重要组成部分。
  3. 计算与通信融合的需求。 传统 host 端集合 API 将通信与计算严格分离,无法支持 kernel 内通信、通信计算重叠等细粒度融合模式。
  4. 与上游生态对齐的需求。 PyTorch 上游(CUDA/NCCL)已建立 MemPool + backend.mem_allocator + register_mem_pool(pool, symm=...) 的标准用法,vLLM 等推理框架的 CUDA Graph / symmetric memory 路径均基于该用法开发。torch_npu 提供同构接口后,上层框架可以最小代价迁移到 NPU。

1.3 目标

  1. 支持 Memory pool 形式的 symmetric memory 使用,对齐 NCCL backend 的接口语义:
    • c10d::Backend::getMemAllocator:返回基于 HcclMemAlloc/HcclMemFree 的 HCCL 专用分配器;
    • ProcessGroupHCCL::registerMemPool(pool, symm):将 MemPool 中已分配的段批量注册到 HCCL 通信域,并通过分配器 trace hook 保证注册后新分配的段自动注册;
    • ProcessGroupHCCL::deregisterMemPool(pool):注销并清理映射关系。
  2. 同时支持对称窗口注册(symm=True)与用户 buffer 注册(symm=False)两种模式。
  3. 保证与 NPU 缓存分配器(NPUCachingAllocator)的 MemPool 机制兼容
  4. 提供多进程(8 卡)与单卡的功能正确性测试。

2. 用例分析

2.1 用例一:模型 MoE 训练四类集合通信(symm=True)

场景描述: DeepSeek 系列 MoE 模型训练中,通信由四类集合原语构成:MoE 层专家并行(EP) 触发 All-to-All;参数在前向/反向计算前通过 AllGather 聚合;反向阶段梯度通过 ReduceScatter 归约并重新分片;梯度同步使用 AllReduce。业务在初始化阶段通过 backend.mem_allocator 创建 MemPool,并使用HCCL 直接使用对称内存做通信。
功能点: HCCL 分配器获取、MemPool 分配路由、对称窗口注册、注册后增量分配自动注册、反注册。
关键性能指标: 1000 step无精度和性能问题,在通信数据量大的场景下有性能提升。

2.2 用例二:八卡4种集合通信测试

场景描述: 模拟 DeepSeek MoE 训练的完整通信画像,在单一用例内覆盖 AllReduce、AllGather、ReduceScatter、All-to-All 四类集合原语。测试以 8 进程 HCCL 通信组运行,先通过 backend.mem_allocator 创建 MemPool,在 use_mem_pool 上下文中为四类集合通信分别分配输入/输出 buffer,register_mem_pool(pool, symm=True) 后依次执行四类集合通信,各原语输出与注册前非池化路径的基线结果逐一比对,最后 deregister_mem_pool 完成反注册。数据类型统一取 float32(symmetric kernel 不支持 int64 等类型),消息规模覆盖 KB~MB 级以贴近小消息高频场景。

功能点: HCCL 分配器获取、MemPool 分配路由、对称窗口注册、四类集合通信(all_reduce / all_gather_into_tensor / reduce_scatter_tensor / all_to_all_single)对池内注册内存的直接消费、注册后增量分配自动注册、反注册。
关键性能指标: 四类集合通信输出与非注册路径基线逐位一致(正确性不回归);四类原语端到端延迟相对传统路径下降;消除 staging buffer 拷贝——AllGather 输出、ReduceScatter 输入、All-to-All dispatch 输入均直接产自/落于池内注册内存,无需中转;注册开销仅在 pool 注册时发生一次,不进入任一集合通信的关键路径;反注册后池内内存安全回收、流程无泄漏。

2.3 约束与限制

npudriver要求:26.1.0
排查方法shell中输入:
$npu-smi info
顶部显示版本号>= 26.1.0

cann版本要求:
cann版本>=9.2.0

接口约束同HcclCommSymWinRegister:

  • 仅支持Atlas A3 训练系列产品/Atlas A3 推理系列产品的超节点内通信。
  • 仅支持通信算子展开模式为AI CPU的场景。
  • 仅支持超节点内AI Server间使用HCCS链路进行SDMA通信的场景,不支持使用RoCE进行RDMA通信的场景(即不支持设置环境变量HCCL_INTER_HCCS_DISABLE为"TRUE",单机场景该环境变量无效)。
  • 仅支持对称组网,即每个Server内卡数相同的场景。
  • 该接口仅支持集合通信算子AllGather、ReduceScatter、AllReduce、AllToAll。
  • 需确保通信域中的所有rank同时调用该注册接口。
  • 所有rank的输入地址映射的物理内存大小一致(对称内存注册按物理内存的大小对齐)。
  • 所有rank调用该接口时,输入的size参数需要保持一致。
  • 使用对称内存功能时,算子的输入、输出内存必须调用此接口注册为对称内存。
  • 调用该接口注册的内存需要使用HcclCommSymWinDeregister接口解注册。

参考:https://www.hiascend.com/document/detail/zh/CANNCommunityEdition/latest/API/hcclug/docs/zh/api_ref/comm_mgr_c/HcclCommSymWinRegister.md

3.1 总体方案

实现 symmetric memory 的 mem_pool 部分,对外提供三个接口(与 NCCL backend 用法同构):

  1. c10d::Backend::getMemAllocator:返回基于 HcclMemAlloc/HcclMemFree 的自定义分配器(经 createCustomAllocator 包装),供 torch_npu.npu.MemPool 使用;
  2. ProcessGroupHCCL::registerMemPool(pool, symm):通过 snapshot 将 pool 存量段批量注册到 HCCL 通信域(symm=TrueHcclCommSymWinRegister 对称窗口,FalseHcclCommRegister 用户 buffer),并挂接分配器 trace hook,使后续新增/释放的段自动注册/反注册;
  3. ProcessGroupHCCL::deregisterMemPool(pool):反注册 pool 全部段并清理映射。

实现上:HCCLComm 新增 registerSegment/deregisterSegment 管理段句柄;hcclCommMemPoolMap 维护通信域与 pool 的映射,随通信域销毁同步清理;新增 HCCL 符号经 TORCH_NPU_LOAD_FUNCTION 动态加载,兼容旧版 CANN。

使用方式:

        # Use HCCL memory allocator
        pool = torch.npu.MemPool(backend.mem_allocator)

        # allocate memory with hccMemAlloc
        with torch.npu.use_mem_pool(pool):
            tensor = torch.arange(1024 * 1024 * 2, device=device)

        # register buffers to HCCL
        backend.register_mem_pool(pool, symm=True)

        pg.allreduce(tensor).wait()
        torch.npu.synchronize(device=device)

        # de-register buffers from HCCL
        backend.deregister_mem_pool(pool)

3.2 技术选型

symmetric memory 已成为业界优化小消息集合通信的主流方向,开源生态中已形成三种典型后端实现:

后端 实现机制
VMM(虚拟内存管理) 基于驱动级虚拟内存管理,使用handle直接读写对端内存,使用put_signal等同步
通信库 (如NCCL) 由通信库分配内存并注册对称窗口,建立跨 rank 映射,注册完后自动走symmetric加速
NVSHMEM 支持NVSHMEM内存注册接口,以及专属算子

开源社区的使用形态分为两种:一种是 tensor 级的 rendezvous API;还有一种是 MemPool 形式MemPool(backend.mem_allocator) + register_mem_pool(pool, symm)),将对称内存接入框架的缓存分配器体系。PyTorch 上游(torch.distributed._symmetric_memory)两种形态均已支持。

NPU 侧,HCCL 已提供 HcclCommSymWinRegister/HcclMemAlloc 等等价底层能力,但 torch_npu 尚未将其接入 MemPool 体系,上层框架的 symmetric memory 路径无法在 NPU 上复用。本方案优先实现ProcessGroupHCCL中提供的symmetric memory功能。

3.3 功能与性能设计

功能: 重复注册直接返回成功;每段记录注册类型,注销时自动匹配对应接口;对齐业界的symmetric memory使用方式。
性能: 注册内存开销集中在 pool 注册或者内存申请时,通信关键路径无额外开销;HCCL 直接操作用户内存,消除 hccl buffer的拷贝。

使用方式:

pool = torch.npu.MemPool(backend.mem_allocator)
backend.register_mem_pool(pool, symm=True)
with torch.npu.use_mem_pool(pool):
    x = torch.empty(1024, device="npu:0")   

3.4 安全隐私与DFX设计

  • 不影响原代码: 全部为新增接口,不使能该特性的用户行为完全不变;
  • 版本和硬件平台要求明确: 依赖 CANN 提供 SymWinRegister/MemAlloc 等符号,旧版本仅在显式调用新接口时报错,不影响其他功能;
  • DFX设计:统计当前正在使用的symmetric memory数量,pta开启debug日志之后可以看到哪些集合通信操作使用了symmetric memory

3.5 编程与调用设计

3.5.1 编程模型基本设计
3.5.2 接口定义与设计

共实现三个用户面接口,依赖 PTA/HCCL 侧的 7 个接口:

用户面接口 依赖
ProcessGroupHCCL::registerMemPool HcclCommSymWinRegister / HcclCommRegisterattachAllocatorTraceTrackerNPUCachingAllocator::snapshot
ProcessGroupHCCL::deregisterMemPool HcclCommSymWinDeregister / HcclCommDeregisterNPUCachingAllocator::snapshot
c10d::Backend::getMemAllocator NPUPluggableAllocator::createCustomAllocatorHcclMemAllocHcclMemFree

注:issue 初稿中写的 ProcessGroupNCCL::*hcclCommWindowRegister 为笔误,以本表(与实际实现一致)为准。

3.5.2.1 backend.mem_allocator(c10d::Backend::getMemAllocator)
  • 接口描述: 返回 HCCL 专用内存分配器(静态共享单例)。该分配器的 alloc/free 分别经 NPUGuard 切换设备后调用 HcclMemAlloc/HcclMemFree,使 MemPool 内分配的内存天然位于 HCCL 可注册的对称堆上。

  • 接口原型:

    std::shared_ptr<c10::Allocator> ProcessGroupHCCL::getMemAllocator() override;
    

    Python 侧为只读属性 backend.mem_allocator

  • 输入/输出参数: 无输入参数。

  • 返回参数: c10::Allocator 智能指针;Python 侧直接传给 torch_npu.npu.MemPool(allocator)

  • 异常处理: 底层 HcclMemAlloc 符号缺失或分配失败时抛出含 NOT_FOUND/HCCL 错误码的异常。

  • 约束说明:

  • 变更说明:

  • 调用参考代码: pool = torch_npu.npu.MemPool(backend.mem_allocator)

3.5.2.2 backend.register_mem_pool(ProcessGroupHCCL::registerMemPool)
  • 接口描述: 将 MemPool 当前已分配的全部段注册到 pool 所在 device 的 HCCL 通信域;同时挂接分配器 trace hook,使注册之后该 pool 内新分配的段自动注册、释放的段自动注销。

  • 接口原型:

    void ProcessGroupHCCL::registerMemPool(c10_npu::MemPool* pool, bool symm = false);
    

    Python 侧:backend.register_mem_pool(pool, symm=True)

  • 输入/输出参数:

    参数名称 输入/输出 类型 描述 取值范围
    pool 输入 c10_npu::MemPool* 待注册的内存池,其 device 决定目标通信域 须由 MemPool(backend.mem_allocator) 创建
    symm 输入 bool true=对称窗口注册(HcclCommSymWinRegister)当前不支持false true
  • 返回参数: 无。

  • 异常处理: 通信域未初始化时报错并提示可通过 init_process_group(device_id=...) 预建。

  • 约束说明: 调用前通信域必须已存在;。当前仅支持非symmetric memory

  • 变更说明: 无。

  • 调用参考代码:

3.5.2.3 backend.deregister_mem_pool(ProcessGroupHCCL::deregisterMemPool)
  • 接口描述: 将 pool 全部已注册段从 HCCL 通信域反注册,并从 hcclCommMemPoolMap 移除映射;反注册后该 pool 的新增段不再自动注册。

  • 接口原型:

    void ProcessGroupHCCL::deregisterMemPool(c10_npu::MemPool* pool);
    

    Python 侧:backend.deregister_mem_pool(pool)

  • 输入/输出参数:

    参数名称 输入/输出 类型 描述 取值范围
    pool 输入 c10_npu::MemPool* 待反注册的内存池 必须已通过 register_mem_pool 注册
  • 返回参数: 无。

  • 异常处理: 反注册未注册的 pool 抛出 "Trying to unregister a pool that was not previously registered";。

  • 约束说明:

  • 变更说明:

3.5.3 编程手册设计

为了帮助开发者能快速上手开发,要设计好本提案相关特性/功能的《编程手册》,要包含哪些内容和章节,单独输出还是共用,在已有的手册中更新还是输出等。确保最后输出的《编程手册》中有相关变更内容。

4. 测试设计

1. 单卡测试

symm=true:
测试单卡情况下使用symmetric mempool接口在all_reduce, all_gather, reduce_scatter和all_to_all 功能正常。
验证Allocator 绑定、MemPool 分配、register/deregister 接口连通 ,等功能是否正常。

symm=false:
symmetric memroy注册不成功,打印异常

2. 多卡测试

symm=true:
跨 rank 把 pool 的段注册进同一通信域、注册后在这些内存上跑 all_reduce, all_gather, reduce_scatter和all_to_all结果正确、device 解析正确、反注册干净,功能正常。
验证Allocator 绑定、MemPool 分配、register/deregister 接口连通 ,等功能是否正常。同时和不使用symmetric memory的集合通信比较在单机8卡的情况下单次集合通信性能。

用例设计:

用例分类 测试类别 核心问题 判定方式
correctiness 精度正确性 symm 路径与 native 路径结果是否一致 单标杆比对(native 为标杆)
perf 性能基准 各消息大小下的耗时与总线带宽 无阈值,输出 median 耗时 + busbw
stress 连续稳定性 长时间反复通信是否泄漏/衰减/挂死 内存增长 <5%、性能衰减 <10%、跑满不挂
edge 边界与异常 极端输入/非法参数下行为是否符合预期 逐用例预期判定(PASS/FAIL/SKIP)

一:correctiness

四个算子的 correctiness 脚本采用完全一致的测试框架:

  • 双通信域并存:启动时一次性建立两个通信域,全程不销毁——默认通信域跑 native 路径(普通显存)得结果 t1;new_group + 注册对称内存池的 symm域跑得结果 t2
  • 确定性输入:固定 seed,保证两个 phase 的输入逐位一致(这是比对成立的前提)
  • 单标杆比对 按照精度标准比较native路径和symmetric memory的精度差异
  • 判定等级:PASS / FAIL / ERROR(symm 不支持但 native 正常)/ SKIP(native 不支持或版本不支持)

通用shape

类型 示例 shape 用途
标量 [1] 边界
1D 小 [1024] 小消息延迟
1D 中 [65536] 中等消息
1D 大 [1048576] 大消息带宽
2D [128, 128] 常规矩阵
3D [12, 56, 256] 多维 tensor
4D [8, 3, 224, 224] 图像类

allreduce测试用shape
使用通用shape测试

allgather测试用shape
使用通用shape+专有shape

专有shape:

类型 input shape output shape (concat形式) 说明
小 1D [1024] [4096] dim0 = 1024 × 4
中 1D [65536] [262144] dim0 = 65536 × 4
2D [128, 256] [512, 256] dim0 = 128 × 4
3D [12, 56, 56] [48, 56, 56] dim0 = 12 × 4

reduce_scatter测试用shape
使用通用shape+专有shape

专有shape:

类型 output shape input shape (concat) 说明
小 1D [1024] [4096] 反向 allgather
中 1D [16384] [65536]
2D [128, 256] [512, 256] dim0 拆分
3D [12, 56, 256] [48, 56, 256] dim0 拆分

alltoall 测试用shape
使用通用shape+专有shape

专有shape:

类型 input shape (all_to_all_single) 各 rank 的 chunk 说明
等长小 [4096] 每 chunk 1024 dim0 整除 world_size
等长大 [1048576] 每 chunk 262144 大消息
不等长 各 rank 各 chunk 不同 all_to_all
2D 等长 [8, 128] 每 chunk [2, 128] dim0 整除 world_size

二:perf

# test/distributed/hccl_validation/<算子名>/test_<算子名>_perf.py
# 核心逻辑:warmup 10 次 → 计时 100 次 → 取中位数 → 算带宽

import time
import torch
import torch.distributed as dist
import torch_npu

def benchmark_<算子名>(rank, world_size, msg_size_bytes, dtype, iterations=100, warmup=10):
    """单次 benchmark"""
    num_elements = msg_size_bytes // dtype.itemsize
    tensor = torch.randn(num_elements, dtype=dtype, device=f'npu:{rank}')
    
    # Warmup
    for _ in range(warmup):
        dist.<算子调用>(tensor)  # 替换为实际算子调用
    torch.npu.synchronize()
    
    # Timed iterations
    start = time.perf_counter()
    for _ in range(iterations):
        dist.<算子调用>(tensor)  # 替换为实际算子调用
    torch.npu.synchronize()
    elapsed = time.perf_counter() - start
    
    median_time_s = elapsed / iterations
    bus_bandwidth_gb_s = (msg_size_bytes * 2 * (world_size - 1) / world_size) / median_time_s / 1e9
    
    return median_time_s, bus_bandwidth_gb_s

性能测试矩阵

算子: <算子名>
─────────────────────────────────────────────────────────────
消息大小:  1KB, 4KB, 16KB, 64KB, 256KB, 1MB, 4MB, 16MB, 64MB, 256MB, 512MB
dtype:     FP32, FP16, BF16,
卡数:      2, 4, 8 (机内); 16(跨机, Phase 2)
迭代次数:  100 (warmup=10) 
─────────────────────────────────────────────────────────────

三:stress

测试目标

  • 1000+ 次迭代无 hang、无 crash
  • 内存不持续增长(leak 检测)
  • 性能不衰减(thermal throttling 检测)

测试脚本关键部分

import gc
import torch_npu

def continuous_stress_test(rank, world_size, iterations=10000):
    """连续压测"""
    tensor = torch.randn(1024 * 1024, dtype=torch.float32, device=f'npu:{rank}')
    
    mem_samples = []
    time_samples = []
    
    for i in range(iterations):
        start = time.perf_counter()
        dist.<算子调用>(tensor)  # 替换为实际算子调用
        torch.npu.synchronize()
        elapsed = time.perf_counter() - start
        
        # 每 100 次采样一次内存和时间
        if i % 100 == 0:
            time_samples.append(elapsed)
            reserved = torch.npu.memory_reserved(rank) / 1024 / 1024
            mem_samples.append(reserved)
            
        if i % 1000 == 0 and rank == 0:
            print(f"iteration {i}/{iterations}, time={elapsed*1000:.2f}ms, "
                  f"memory={reserved:.2f}MB")

    # 分析趋势
    # 1. 内存最后 10 次采样 vs 前 10 次:增幅应 < 5%
    # 2. 耗时最后 10 次采样 vs 前 10 次:增幅应 < 10%

判定标准

指标 标准 不通过时排查方向
内存增长 < 5% (10000 iters) HCCL buffer 泄漏、completion queue 未释放
性能衰减 < 10% (10000 iters) thermal throttling、buffer 碎片
无 hang 100% 通过 deadlock、同步原语问题
无 crash 100% 通过 segfault、OOM

edge边界测试

通用边界用例

序号 测试场景 输入 预期行为
1 空 tensor torch.empty(0) 正常返回,不报错(或返回明确错误)
2 标量 tensor torch.tensor([42.0]) 正确通信
3 1-element torch.randn(1) 正确通信
4 大 tensor (>2GB) torch.randn(536870912) (FP32 → 2GB) 正确通信,无 OOM
5 非连续 tensor torch.randn(100, 100).T.contiguous()[::2] 正确通信或报错
6 stride != shape torch.randn(100, 100).T 正确通信或报错
7 输入在错误 device torch.randn(100, device='cpu') 明确报错
8 输入在错误 dtype 要求的 dtype 不匹配 明确报错
9 world_size=1 单 rank 正常返回(NOOP)

5. 缺点和风险

  1. 对称内存的窗口大小hcclSymWinMaxMemSizePerRank默认值为16G,可能不够大,需要根据实际物理机内存情况配置
  2. 当前不支持传入symm=false的参数用于registerMempool

6. 现有技术

7. 未解决问题

当前reduce_scatter可能存在精度问题,待定位

附录

  • 参考资料链接
  • 术语表
  • 文档更新计划

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
ascend-robotascend-robot成员
6月22日 添加了label:rfc
Zzzzzzzk
6月22日 修改了issue 的描述
Zzzzzzzk
6月24日 修改了issue 的描述
Zzzzzzzk
6月24日 关联了pull request:Make HCCL mem pool registration look up comm without creating it
Zzzzzzzk
6月24日 修改了issue 的描述
此处折叠了65条消息 查看更多
Zzzzzzzk
15 天前 关联了pull request:Add support for hccl memory pool symmetric memory support
Xxuyun15成员
8 天前 关联了pull request:[feat][1/n] mempool基本结构对齐pytorch: 仅回合allocate重构,不删除mempoolcontext(API一致但是ABI不兼容)
Xxuyun15成员
5 天前 关联了pull request:[feat][1/n] mempool基本结构对齐pytorch: 仅回合allocate重构,不删除mempoolcontext(API一致但是ABI不一致)
Xxuyun15成员
5 天前 关联了pull request:[feat][1/n] mempool基本结构对齐pytorch: 仅回合allocate重构,不删除mempoolcontext(API一致但是ABI不一致)
Xxuyun15成员
5 天前 关联了pull request:[feat][1/n] mempool基本结构对齐pytorch: 仅回合allocate重构,不删除mempoolcontext(API一致但是ABI不一致)