news 2026/9/9 22:43:15

如何用 PyTorch Symmetric Memory 让自定义 CUDA 内核直接读写对端 GPU 显存?

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
如何用 PyTorch Symmetric Memory 让自定义 CUDA 内核直接读写对端 GPU 显存?

如何用 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)

文档对这一步的强调:

  • emptyrendezvous必须在组内所有 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_ptrssignal_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]);offsetdstdtype 的元素为单位,默认0dst可以是普通 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 包入口:emptyrendezvousset_backendget等 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/9 22:42:55

Java对接微信商家转账到零钱:接口选型、签名与回调避坑指南

简介&#xff1a;面向Java开发者的微信企业转账到零钱功能实现资料&#xff0c;聚焦企业付款、工资奖金发放、退款等典型业务场景。资源包内含两个核心Java文件&#xff0c;一个用于生成请求签名&#xff0c;另一个封装转账接口调用与参数组装&#xff0c;可直接借鉴到Spring等…

作者头像 李华
网站建设 2026/9/9 22:41:46

STM32 PWM呼吸灯实战:TIM3配置与引脚重映射详解

简介&#xff1a;这是一份基于STM32F1系列HAL库的双极性SPWM波形生成工程代码包&#xff0c;面向嵌入式开发、电力电子与电机控制领域的学习者和工程师&#xff0c;可用在逆变器、电机驱动等场景中产生逼近正弦波的调制信号&#xff0c;并支持通过修改滤波器参数改变输出频率。…

作者头像 李华
网站建设 2026/9/9 22:39:16

ImageNet按需下载:构建自定义数据集的轻量方案

简介&#xff1a;面向图像分类与计算机视觉研究者&#xff0c;提供一套基于 Python 3 的 ImageNet 子集自助下载方案。核心脚本可指定类别数量和每类图片张数&#xff0c;自动从 ImageNet 图像 URL 中随机筛选并抓取样本&#xff0c;用于快速搭建训练集、验证集或进行小规模实验…

作者头像 李华
网站建设 2026/9/9 22:38:03

办公设备效率评估:从卡顿诊断到软硬件替换的实操指南

你是否曾经被一台“性能充沛”却日夜卡顿的办公电脑折磨到崩溃?明明每天都在赶进度&#xff0c;却被软件启动速度、文件加载延迟这些看似微小的问题不断打断思路。从我的实际体验来看&#xff0c;办公设备的效率评估绝不只是“跑个分”“看个参数”那么简单&#xff0c;它更像…

作者头像 李华
网站建设 2026/9/9 22:35:01

2026时序数据库选型:金仓融合多模架构如何破解双库之痛

从2025年下半年开始&#xff0c;我陆续接到好几个项目团队的同样诉求&#xff1a;原本只用关系型数据库做业务系统&#xff0c;现在因为设备数据、车联网轨迹、能源计量这类时序数据暴涨&#xff0c;被迫在架构里引入新的时序数据库。可引进来之后麻烦更多了——两套库、两套账…

作者头像 李华
网站建设 2026/9/9 22:34:26

2026公众号投票活动搭建教程:3分钟创建+推文嵌入全流程

做公众号运营的朋友应该都有体会&#xff1a;想在文章里加一个投票互动&#xff0c;看似简单&#xff0c;实际操作起来却常常碰壁。公众号自带的投票功能最多支持30个选项&#xff0c;不能展示图片视频&#xff0c;不防刷&#xff0c;也不能导出数据。想办一场像样的评选活动&a…

作者头像 李华