如何用 PyTorch Symmetric Memory 让自定义 CUDA 内核直接读写对端 GPU 显存?
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
在多 GPU 训练中,高层 collective API 往往覆盖不了自定义通信模式的延迟要求,你需要自己写内核,在内核里直接读写对端 GPU 的显存。PyTorch 的torch.distributed._symmetric_memory(下称 SymmMem)提供这条路:每个 rank 创建对称张量,通过 rendezvous 交换句柄后,内核即可拿到对端 buffer 的本地地址和用于同步的信号垫(signal pad),像访问本地数据一样访问对端内存。文档明确标注该包目前处于 alpha 阶段、API 可能变动,本文的操作路径均出自该文档。
前置条件:硬件、进程组与后端
- 硬件背景:文档以高带宽互连(NVLink、InfiniBand 或 RoCE)让对端 GPU 全局显存可直接访问的系统为前提。若选择
NCCL后端,文档额外要求 NCCL 2.27 及以上,且处于单一 NVLink 域(每个 rank 都能经直连 NVLink 到达)。 - 进程组:先完成
dist.init_process_group(),每个进程运行在自己的 GPU 上。 - 后端:用
symm_mem.set_backend(name)选择对称内存后端,目前支持"NVSHMEM"、"CUDA"、"NCCL"。这是全局设置,影响之后所有symm_mem.empty()调用,且一旦分配过第一个对称内存张量就不能再切换。文档的基础示例未显式设置后端;需要NCCL时按示例显式调用symm_mem.set_backend("NCCL")。
第 1 步:创建对称张量并完成 rendezvous
import torch.distributed as dist import torch.distributed._symmetric_memory as symm_mem dist.init_process_group() rank = dist.get_rank() # Allocate a tensor t = symm_mem.empty(4096, device=f"cuda:{rank}") # Establish symmetric memory and obtain the handle hdl = symm_mem.rendezvous(t, dist.group.WORLD)文档对这一步的强调:
empty和rendezvous必须在组内所有 rank 上以相同顺序调用。rendezvous是集合操作,且是 host 阻塞的初始化操作:首次调用要在进程间完成句柄交换与映射,并同步 host 与设备。它无法被排到 CUDA stream 上,也不能被 CUDA graph 捕获。文档建议初始化时分配一次对称 buffer、之后复用返回的 handle;对同一张量再次调用会返回缓存的 handle。group参数可以传组名字符串或ProcessGroup对象;张量的 shape、dtype 与 device 类型在所有参与进程上必须一致。
第 2 步:把对端地址与信号垫传给内核
拿到 handle 后,以下三项可直接传给内核:
hdl.buffer_ptrs # 各 peer 上的对称 buffer 地址 hdl.multicast_ptr # 多播指针(硬件支持时可用) hdl.signal_pad_ptrs # 用于同步的信号垫地址文档说明:buffer_ptrs指向的数据可以像普通本地数据一样访问,并建议像本地数据一样使用向量化访问来提高效率。SymmMem 提供的同步原语与 CUDA Graph 兼容,操作对象就是每次对称内存分配所附带的信号垫。文档指出内核既可以用 CUDA 写,也可以用 Triton 写,机制相同:内核接收 handle 提供的对端地址与信号垫地址,在内核里直接读写对端显存。
第 3 步:文档给出的完整内核示例(Triton one-shot all-reduce)
@triton.jit def one_shot_all_reduce_kernel( buf_tuple, signal_pad_ptrs, output_ptr, numel: tl.constexpr, rank: tl.constexpr, world_size: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): ptx_utils.symm_mem_sync( signal_pad_ptrs, None, rank, world_size, hasSubsequenceMemAccess=True ) pid = tl.program_id(axis=0) block_start = pid * BLOCK_SIZE while block_start < numel: offsets = block_start + tl.arange(0, BLOCK_SIZE) mask = offsets < numel acc = tl.zeros((BLOCK_SIZE,), dtype=tl.bfloat16) for i in tl.static_range(world_size): buffer_rank = buf_tuple[i] x = tl.load(buffer_rank + offsets, mask=mask) acc += x tl.store(output_ptr + offsets, acc, mask=mask) block_start += tl.num_programs(axis=0) * BLOCK_SIZE ptx_utils.symm_mem_sync( signal_pad_ptrs, None, rank, world_size, hasPreviousMemAccess=True )buf_tuple对应hdl.buffer_ptrs,signal_pad_ptrs对应hdl.signal_pad_ptrs。- 同步工具模块
ptx_utils的实现不包含在文档中;文档指引到 kraken 项目(meta-pytorch/kraken)查看完整 utilities 与常见模式示例。 - 文档未给出内核的 launch 代码;从内核体内的步长
tl.num_programs(axis=0) * BLOCK_SIZE可以判断,grid 启动需使“程序数 × BLOCK_SIZE”覆盖numel。
文档说明,内核开头与结尾的两次symm_mem_sync保证所有进程看到一致的数据——这是文档给出的此类内核的一致性保证依据。
结果验证:文档提供的判断方式
- 内核路径:文档没有提供独立的检查脚本,其依据是上述两次同步的语义——前置同步确保对端已就绪(
hasSubsequenceMemAccess=True),后置同步在写入完成后通知对端(hasPreviousMemAccess=True),从而“all the processes see consistent data”。 - NCCL 后端路径:文档给出可直接执行的日志检查。选择
NCCL后端后,rendezvous会把分配做 window 注册,标准集合通信(dist.all_reduce等)即可被分派到 NCCL 的对称内核:
symm_mem.set_backend("NCCL") x = symm_mem.empty(1024 * 1024, dtype=torch.bfloat16, device=device) symm_mem.rendezvous(x, group=dist.group.WORLD.group_name) dist.all_reduce(x, op=dist.ReduceOp.SUM)用下面的命令检查内核名(train.py为文档示例中的入口脚本,替换为你自己的入口):
NCCL_DEBUG=INFO NCCL_DEBUG_SUBSYS=TUNING python train.py文档示例输出(示例结果,不是必须得到的固定值):
AllReduce [Symmetric]: 2097152 Bytes -> Kernel AllReduce_RSxLDMC_AGxSTMC nchannels 16 nthreads 512 nWorks 1在 profiler 里,设备内核名应形如ncclSymkDevKernel_*(例如ncclSymkDevKernel_AllReduce_AGxLLMC_R_sum_bf16),而普通 NCCL 路径是ncclDevKernel_*。
可选路径:先跑内置 op,确认环境就绪
文档的基础示例展示了在同一对称张量上直接调用内置 op:
# Most SymmMem ops are under the torch.ops.symm_mem namespace torch.ops.symm_mem.one_shot_all_reduce(t, "sum", group)两个文档明确给出的注意点:
torch.ops.symm_mem是 “op namespace” 而非 Python 模块,不能import torch.ops.symm_mem,也不能from torch.ops.symm_mem import one_shot_all_reduce,直接按上例调用即可。reduce_op目前只支持"sum",第三个参数是组名(字符串)。- 同一组的所有 symm_mem 集合通信必须从同一个 CUDA stream 发起。内核通过以 block ID 索引的共享信号垫同步 rank,没有 per-stream 隔离;从不同 stream 并发发起同一组的集合通信会死锁。确需多流时,用
stream.wait_stream()/current_stream.wait_stream()串行到一条专用 stream 上。
可选路径:不写内核也能读对端内存——one-sided get
如果目标只是把对端对称内存的一段读进来,文档提供了 host 侧getAPI:
src = symm_mem.empty(1024, device=device) hdl = symm_mem.rendezvous(src, group) if dist.get_rank(group) == 0: dst = torch.empty((512,), device=device) # Copy the last 512 elements of the peer's allocation into dst. symm_mem.get(dst, hdl, peer=1, offset=512)文档明确的语义:拷贝的元素数由dst推断,想只拷一部分就传视图(如dst[:n]);offset以dstdtype 的元素为单位,默认0;dst可以是普通 CUDA 张量或另一个对称张量,必须与hdl同设备且由连续内存支撑;拷贝在当前 CUDA stream 上发起。
可选分支:跨节点与大规模 rendezvous
- 跨节点:文档说明多节点访问依赖支持 RDMA 的 NIC,PyTorch 提供 NVSHMEM 插件扩展 Triton 内核的跨节点能力,可在内核里发起 put:
import torch.distributed._symmetric_memory._nvshmem_triton as nvshmem from torch.distributed._symmetric_memory._nvshmem_triton import requires_nvshmem @requires_nvshmem @triton.jit def my_put_kernel( dest, src, nelems, pe, ): nvshmem.put(dest, src, nelems, pe)requires_nvshmem装饰器声明内核依赖 NVSHMEM device 库;Triton 编译时会在系统路径中搜索该库,找到则包含必要的 device assembly。
- 大规模 rendezvous:默认
rendezvous经 TCPStore 交换元数据,文档给出的容量参考是约 20 万 QPS;在 10k 总 rank、72-rank NVLink 组的例子里单次 rendezvous 约 3.6s,10 万 rank 时增长到约 36s。改用进程组的 NCCL allgather,需在进程组选项中设置use_pg_for_symm_mem_rendezvous。若进程组只用于对称内存、之后不再做普通集合通信(例如专家并行组),rendezvous 后可abort()释放 NCCL communicator——handle 只依赖已映射的内存,仍然可用:
opts = dist.ProcessGroupNCCL.Options() opts.use_pg_for_symm_mem_rendezvous = True ep_pg = dist.new_group(ep_ranks, pg_options=opts) t = symm_mem.empty(size, device=device) hdl = symm_mem.rendezvous(t, group=ep_pg) # Release the NCCL communicator since ep_pg won't be used for collectives. # The symm_mem handle is still usable — it only needs the mapped memory. ep_pg.abort()注意:启用use_pg_for_symm_mem_rendezvous时,若进程组的 NCCL communicator 尚不存在会被惰性创建。
限制与边界
- API 处于 alpha 阶段,签名可能变动。
- 后端在首次分配对称张量后不可更改。
rendezvous是 host 阻塞的初始化操作,不能排入 CUDA stream、不能被 CUDA graph 捕获;应初始化一次、复用 handle。- 同一组的 symm_mem 集合通信必须单 stream 发起,否则死锁。
multimem_*系列 op 额外要求硬件多播支持(NVIDIA 上需要 NVLink SHARP),与本文指针访问路径不同,按需使用。
延伸阅读
- Symmetric Memory 文档:本文全部示例的出处,含 Memory Pool、NCCL Symmetric Kernels、Copy Engine Collectives 等进阶章节
- symm_mem 包入口:
empty、rendezvous、set_backend、get等 API 的完整 docstring - NVSHMEM Triton 扩展:跨节点 put/get 的 Triton 端实现
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考