十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

如何用 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 步创建对称张量并完成 rendezvousimport 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, devicefcuda:{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-reducetriton.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, hasSubsequenceMemAccessTrue ) pid tl.program_id(axis0) 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,), dtypetl.bfloat16) for i in tl.static_range(world_size): buffer_rank buf_tuple[i] x tl.load(buffer_rank offsets, maskmask) acc x tl.store(output_ptr offsets, acc, maskmask) block_start tl.num_programs(axis0) * BLOCK_SIZE ptx_utils.symm_mem_sync( signal_pad_ptrs, None, rank, world_size, hasPreviousMemAccessTrue )buf_tuple对应hdl.buffer_ptrssignal_pad_ptrs对应hdl.signal_pad_ptrs。同步工具模块ptx_utils的实现不包含在文档中文档指引到 kraken 项目meta-pytorch/kraken查看完整 utilities 与常见模式示例。文档未给出内核的 launch 代码从内核体内的步长tl.num_programs(axis0) * BLOCK_SIZE可以判断grid 启动需使“程序数 × BLOCK_SIZE”覆盖numel。文档说明内核开头与结尾的两次symm_mem_sync保证所有进程看到一致的数据——这是文档给出的此类内核的一致性保证依据。结果验证文档提供的判断方式内核路径文档没有提供独立的检查脚本其依据是上述两次同步的语义——前置同步确保对端已就绪hasSubsequenceMemAccessTrue后置同步在写入完成后通知对端hasPreviousMemAccessTrue从而“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, dtypetorch.bfloat16, devicedevice) symm_mem.rendezvous(x, groupdist.group.WORLD.group_name) dist.all_reduce(x, opdist.ReduceOp.SUM)用下面的命令检查内核名train.py为文档示例中的入口脚本替换为你自己的入口NCCL_DEBUGINFO NCCL_DEBUG_SUBSYSTUNING 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 侧getAPIsrc symm_mem.empty(1024, devicedevice) hdl symm_mem.rendezvous(src, group) if dist.get_rank(group) 0: dst torch.empty((512,), devicedevice) # Copy the last 512 elements of the peers allocation into dst. symm_mem.get(dst, hdl, peer1, offset512)文档明确的语义拷贝的元素数由dst推断想只拷一部分就传视图如dst[:n]offset以dstdtype 的元素为单位默认0dst可以是普通 CUDA 张量或另一个对称张量必须与hdl同设备且由连续内存支撑拷贝在当前 CUDA stream 上发起。可选分支跨节点与大规模 rendezvous跨节点文档说明多节点访问依赖支持 RDMA 的 NICPyTorch 提供 NVSHMEM 插件扩展 Triton 内核的跨节点能力可在内核里发起 putimport 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.6s10 万 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_optionsopts) t symm_mem.empty(size, devicedevice) hdl symm_mem.rendezvous(t, groupep_pg) # Release the NCCL communicator since ep_pg wont 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 的完整 docstringNVSHMEM Triton 扩展跨节点 put/get 的 Triton 端实现【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表