乐于分享
好东西不私藏

AI Infra SHMEM专题 PyTorch Symmetric Memory 实现

AI Infra SHMEM专题 PyTorch Symmetric Memory 实现

系列四:PyTorch Symmetric Memory——从 Tensor 分配到远端 handle

上一篇:系列三:CANN SHMEM 实现
下一篇:系列五:torch_npu NPUSHMEM

本文目标

本文解释 PyTorch 如何把“对称内存”接入 Tensor 与 dispatcher。读完后,你应该能回答:

  • • symm_mem.empty() 和普通 torch.empty() 的根本差异是什么;
  • • allocator、Tensor、Storage、rendezvous、SymmetricMemory handle 分别负责什么;
  • • 为什么 allocation 本身不是 collective,而 rendezvous 是 collective;
  • • remote buffer、signal pad、group、MemPool 和 stream 如何关联;
  • • PyTorch fused all-gather matmul 等功能为什么需要 symmetric memory;
  • • 哪些行为属于 PyTorch v2.11 公共实现,哪些仍是 CUDA-specific 或 backend-specific。

术语速查

名词
本文含义
allocator
为某种 device type 提供 symmetric allocation/free/rendezvous 的后端对象。
empty_strided_p2p
C++ 分配入口;返回可继续 rendezvous 的 Tensor。
rendezvous
各参与 rank 建立 local allocation 与 peer buffers 关联的集体过程。
handle
rendezvous 返回的 SymmetricMemory 抽象对象,暴露 peer buffer、signal pad 和同步能力。
Storage
Tensor 持有底层数据指针和释放逻辑的对象。
TensorImpl
Tensor 的 shape、stride、dtype、device 等实现层元数据。
signal pad
与数据 buffer 分开的 P2P 可访问同步区。
MemPool
把普通 Tensor 分配临时路由到指定 allocator 的内存池机制。
persistent allocation
通过 alloc_id 复用确定地址的内部能力,服务内存规划和编译场景。

源码快照

仓库:https://github.com/pytorch/pytorch
标签:v2.11.0
提交:70d99e998b4955e0049d13a98d77ae1b14db1f45

核心文件:

torch/distributed/_symmetric_memory/__init__.py
torch/csrc/distributed/c10d/symm_mem/SymmetricMemory.hpp
torch/csrc/distributed/c10d/symm_mem/SymmetricMemory.cpp
torch/csrc/distributed/c10d/symm_mem/CUDASymmetricMemory.*
torch/csrc/distributed/c10d/symm_mem/NVSHMEM*
test/distributed/test_symmetric_memory.py

torch.distributed._symmetric_memory 的下划线表明它仍是内部/实验性质较强的模块。读源码时不要仅凭函数存在就推断所有后端都有相同支持。

1. 快速预览

PyTorch symmetric memory 把底层 SHMEM/P2P 能力拆成两步:

empty:      先得到由 symmetric allocator 支撑的本地 Tensor
rendezvous: 再让一组 rank 把各自 local allocation 关联成 peer-accessible handle

完整对象关系:

PyTorch symmetric memory 从本地 Tensor 到通信 handle 的对象链路

图 1:empty 先形成由对称 allocator 支撑的本地 Tensor;rendezvous 再建立 peer-accessible handle。Tensor 与 handle 相关,但不是同一个对象。

这能解释两个常见疑问:

  • • Tensor 是计算框架里的数据对象;
  • • handle 是通信关系对象。

它们相关,但不是同一个东西。

2. Python 入口:empty

公开 Python 实现在:

torch/distributed/_symmetric_memory/__init__.py

empty() 负责:

  1. 1. 规范化 size;
  2. 2. 补默认 dtype 和 device;
  3. 3. 计算 contiguous stride;
  4. 4. 选择 implicit MemPool 或直接 empty_strided_p2p;
  5. 5. 返回 Tensor。

v2.11.0 的真实分支是:

if implicit_pool_enabled and device.type == "cuda":
    mempool = get_mem_pool(device)
    with torch.cuda.use_mem_pool(mempool):
        _SymmetricMemory.empty_strided_p2p(...)
else:
    _SymmetricMemory.empty_strided_p2p(...)

关键边界:

  • • implicit MemPool 默认开启,但这里只对 device.type == "cuda" 生效;
  • • 源码 TODO 明确说,只有把 use_mem_pool 提升为通用 accelerator API 后,这条路径才可设备无关;
  • • PrivateUse1/NPU 在该快照中走 else,不能写成“PyTorch v2.11 已通过通用 MemPool 自动完成 NPU 分配”。

环境变量 TORCH_SYMMMEM_IMPLICIT_POOL=0 可以关闭 CUDA implicit pool 选择。它不改变后端是否注册 allocator。

3. C++ allocator 抽象

SymmetricMemoryAllocator 定义在:

torch/csrc/distributed/c10d/symm_mem/SymmetricMemory.hpp

核心虚函数:

virtual void* alloc(size_t size, int device_idx,
                    const
 std::optional<std::string>& group_name)
= 0;
virtual void free(void* ptr)
= 0;
virtual size_t get_alloc_size(void* ptr)
= 0;
virtual c10::intrusive_ptr<SymmetricMemory> rendezvous(
    void
* ptr, const std::optional<std::string>& group_name)
= 0;
virtual bool has_multicast_support(int device_idx)
= 0;
virtual c10::DeviceType supported_device_type()
= 0;
virtual std::string name()
= 0;

注册表以 DeviceType 找 allocator:

register_allocator(device_type, allocator)
  → AllocatorMap[device_type] = allocator

empty_strided_p2p(device)
  → get_allocator(device.type())
  → allocator->alloc(...)

因此第三方后端的接入点不是修改 Python empty() 中的每个调用者,而是为自己的 device type 注册实现。

4. empty_strided_p2p 如何形成 Tensor

v2.11.0 的 C++ 主链路:

empty_strided_p2p(size, stride, dtype, device, group_name, alloc_id)
  → 计算 numel × element_size
  → get_allocator(device.type())
  → allocator->alloc(bytes, device.index(), group_name)
  → at::from_blob(dev_ptr, size, stride, deleter, options)
  → 返回 Tensor

deleter 捕获 allocator:

Tensor/Storage 生命周期结束
  → deleter(ptr)
  → allocator->free(ptr)

4.1 Tensor、Storage、TensorImpl

可以用三层理解:

Tensor: 用户看到的句柄
  └─ TensorImpl: shape / stride / dtype / device 等
       └─ Storage: data pointer / bytes / deleter / ownership

at::from_blob 把已有指针包装为 Tensor。它不会重新分配数据,但会构造 Tensor/Storage 对象并安装释放回调。

这条通用包装路径要求设备后端确认:它生成的 Storage/TensorImpl 是否携带该后端后续算子所需的全部私有元数据。PyTorch 抽象只规定 symmetric allocator 合同,不会替每个 PrivateUse1 后端补齐其私有 storage 描述。

4.2 allocation 不是 collective

头文件注释明确指出 empty_strided_p2p() 本身不是 collective。原因是:

  • • 每个 rank 可以先完成本地 allocation;
  • • 真正建立 peer association 的步骤是 rendezvous;
  • • 但应用仍必须让各参与 rank 以兼容的 shape、dtype、device 和分配顺序进入后续 rendezvous。

“非 collective API”不代表可以让不同 rank 随意分配不同大小后再 rendezvous。

5. rendezvous:把本地 Tensor 变成通信关系

Python 入口:

symm_mem.rendezvous(tensor, group)

它接受 group name 或 ProcessGroup,后者会转换为 group.group_name。

C++ 主链路:

rendezvous(tensor, group_name)
  → get_allocator(tensor.device().type())
  → tensor.storage().data_ptr().get()
  → allocator->rendezvous(ptr, group_name)
  → intrusive_ptr<SymmetricMemory>

rendezvous 具有 collective 语义:所有参与者需要共同建立映射。后端可能通过 Store bootstrap,也可能使用更高效的控制面;PyTorch 只提供 set_group_info() / get_group_info() 等公共信息。

同一个 allocation 与 group 的 rendezvous 通常只建立一次,后端可缓存 handle。

6. SymmetricMemory handle 提供什么

抽象类同样定义在 SymmetricMemory.hpp。能力分为五组:

能力
典型方法
用途
peer 数据区
get_buffer_ptrs()
、get_buffer()、get_remote_tensor()
访问某个 rank 的对应 buffer。
device pointer table
get_buffer_ptrs_dev()
custom kernel 在设备侧按 rank 查远端 pointer。
signal pad
get_signal_pad_ptrs()
、get_signal_pad()
barrier、signal 或自定义同步。
同步
barrier()
、put_signal()、wait_signal()
建立执行与数据消费顺序。
拓扑/能力
get_rank()
、get_world_size()、has_multicast_support()
算法选择和 peer 定位。

handle 是抽象基类;CUDA、NVSHMEM、NCCL 或 PrivateUse1 后端返回不同派生实现。接口存在不代表某个派生类一定完整支持,后端可以对未实现方法抛出错误。

7. Signal pad 与 channel

signal pad 是与数据 buffer 分开的 P2P 可访问区。它用于同步,不应和用户 payload 随意重叠。

PyTorch 文档强调两个约束:

  • • 成功完成同步后,signal pad 应恢复为全零,以便复用;
  • • channel 用于隔离不同 stream 上的同步,避免两个 barrier 互相消费对方的 signal。

例如:

stream A: barrier(channel=0)
stream B: barrier(channel=1)

如果两者共享同一 channel,A/B 的 signal 可能被错误匹配。

set_signal_pad_size() 必须在任何 symmetric allocation 前调用。大小应随 block 数和 world size 规划,而不是把某个 CUDA 默认常量照搬给其他后端。

8. MemPool 为什么出现

直接 empty_strided_p2p → allocator->alloc → from_blob 能构造 symmetric Tensor,但普通后端往往希望继续走自己的标准 Tensor factory,以保留完整 storage 元数据。

MemPool 的思路是:

注册一个 symmetric raw allocator
  → 创建 no_split 的专用 pool
  → 在 pool context 内调用普通 tensor factory
  → 数据来自 symmetric allocator
  → Tensor 元数据仍由设备后端标准路径创建

v2.11.0 提供:

register_mempool_allocator(device_type, allocator)
get_mempool_allocator(device)
get_mem_pool(device)

get_mem_pool() 预设:

  • • use_on_oom=False:不把 symmetric pool 借给普通 OOM 回退,避免各 rank allocation 状态失配;
  • • no_split=True:不让多个 Tensor 共享一个带 signal pad 的 segment,避免并发同步冲突。

但该版本 Python helper 返回 torch.cuda.MemPool,implicit context 也调用 torch.cuda.use_mem_pool。所以它体现了目标架构,不能泛化为所有 accelerator 已完全对齐。

9. Persistent allocation

内部 C++ empty_strided_p2p 支持 alloc_id:

相同 alloc_id
  → 复用可预测的 allocation 地址
  → 编译器/内存规划可复用已经 rendezvous 的通信 buffer

安全约束是:旧 allocation 的 Storage 仍存活时,不能用同一 alloc_id 再建立冲突的活跃对象。

这是编译和通信 workspace 的高级能力,不应与普通 Python empty(shape) 的基本用法混为一谈。

10. 上层应用为什么需要它

v2.11.0 的 Python 模块包含多类算法:

  • • pipelined all-gather and consume;
  • • fused all-gather matmul / scaled matmul;
  • • fused matmul reduce-scatter;
  • • low-contention all-gather / reduce-scatter;
  • • all-to-all 相关路径;
  • • multicast 优化选择。

共同结构是:

对称 workspace / input
  → 获取 peer buffer 或 remote tensor
  → 用 signal pad 协调 producer-consumer
  → 让通信块与 matmul/consume 块流水化

这比“先完整 all-gather,再启动 matmul”更有机会隐藏通信延迟。但某个 fused op 是否支持 NPU,取决于 dispatcher 注册、后端 handle 能力、kernel 和同步实现,不能从 Python 函数存在直接推断。

11. 支持边界

结论
v2.11.0 状态
empty
 / rendezvous Python API 存在
是
allocator 可按 DeviceType 注册
是
handle 抽象包含 peer buffer/signal/barrier
是
implicit MemPool 对所有 accelerator 通用
否,只在 Python 中对 CUDA 分支启用
每个后端都实现全部 handle 方法
否
Python fused op 对所有设备可用
否,取决于 dispatcher 和 backend
allocation 是 collective
否
rendezvous 是 collective
是

12. 常见误区

  1. 1. empty() 返回后已经能访问远端。
    还没有;必须经过 group rendezvous。
  2. 2. Tensor 本身就是 SymmetricMemory。
    Tensor 持有 local data;handle 描述 peer association 和同步资源。
  3. 3. 注册 allocator 就自动支持所有 fused op。
    fused op 还依赖 remote buffer、signal、kernel、dispatcher 和设备能力。
  4. 4. 有 MemPool API 就说明 NPU 已走 implicit MemPool。
    v2.11 Python 分支仍明确限制在 CUDA。
  5. 5. Signal pad 是额外的数据容量。
    它是同步资源,应遵守清零、channel 和并发使用约束。

13. 读完自检

  • • 为什么 empty_strided_p2p() 可以不是 collective,而 rendezvous() 必须是?
  • • allocator、Storage deleter 和 Tensor 生命周期如何连接?
  • • get_buffer_ptrs() 与 get_buffer_ptrs_dev() 分别服务谁?
  • • signal channel 解决了什么并发问题?
  • • v2.11 为什么不能直接宣称 implicit MemPool 已设备无关?

参考资料

  • • PyTorch v2.11.0 symmetric memory(https://github.com/pytorch/pytorch/blob/v2.11.0)

) 

相关学习资料