
TensorRT 多设备注意力推理实战基于 Context Parallelism 与 MPI/NCCL 的分片部署指南【免费下载链接】TensorRTNVIDIA® TensorRT™ is an SDK for high-performance deep learning inference on NVIDIA GPUs. This repository contains the open source components of TensorRT.项目地址: https://gitcode.com/GitHub_Trending/tens/TensorRT本篇技术指南围绕 TensorRT 官方示例 samples/python/attention_mdtrt 展开完整讲解如何利用 TensorRT 的多设备推理Multi-Device Inference特性将单个自注意力模型切分到多张 GPU 上协同推理。读者将掌握理解张量并行TP与上下文并行CP的区别、编写 sharding 提示文件hint.json、使用 Polygraphy 的multi-device shard工具在 ONNX 图中自动插入 DistCollective 集合通信算子、以及基于 MPI NCCL 编写多进程多 GPU 推理程序。引言为什么需要多设备推理TensorRT 支持将单个模型拆分到多张 GPU 上进行推理。这种能力在两类场景下尤为关键模型过大当模型权重或中间激活超出单卡显存容量时多设备拆分可以突破单卡显存上限降低延迟将计算并行分摊到多个设备上缩短单次推理的端到端时间。TensorRT 为多设备推理提供了多种并行策略主要包括张量并行Tensor Parallelism, TP和上下文并行Context Parallelism, CP。本示例聚焦于上下文并行CP输入序列沿序列维度sequence dimension被切分到多张 GPU 上每个 GPU 独立处理自己负责的那一段序列。由于注意力机制要求每个 token 都必须 attend 到所有其他 token线性投影、归一化等大部分算子可以本地完成但注意力部分必然需要跨设备通信。TensorRT 通过内嵌在模型中的集合通信算子collective operations透明地处理这部分通信。关于上下文并行的原理性细节可参阅论文Context Parallelism for Scalable Million-Token Inference该论文讨论了上下文并行在百万级 token 长序列推理中的可扩展性。多设备模型是如何工作的要在多张 GPU 上运行模型模型必须经过sharding分片把模型拆成每张 GPU 可以独立执行的若干部分。对于上下文并行sharding 意味着将输入序列分给各 GPU并插入通信算子使每张 GPU 仍能对完整序列计算出正确的注意力结果。当一个模型被 sharding 为多设备执行时特殊算子DistCollective会被插入到 ONNX 图中。从 Polygraphy 的 sharding 工具源码shard.py可以看到这些节点以DistCollective为 op_type、命名为DistCollective_N的形式直接注入 ONNX 图ReduceScatter将输入数据跨 GPU 拆分并做归约。用于推理开始时把输入序列分发到各设备AllGather从所有 GPU 收集数据拼成完整张量。在注意力之前用于让每个 GPU 都能看到完整的 K/V 张量在输出投影之后用于重建完整输出。这两个算子不会出现在单设备SDONNX 模型中它们是由polygraphy multi-device shard依据 sharding 提示hints自动插入的。hints 描述了模型应该如何被拆分工具读取单设备 ONNX 图、定位 hints 中指定的注意力层与 I/O 张量然后在恰当的位置插入 DistCollective 算子。hint.jsonsharding 提示文件详解hint.json告诉 Polygraphy 的 sharding 工具如何对模型进行切分。示例仓库中自带的 hint.json 完整内容如下{ parallelism: CP, attention_layers: [ { q: q_scaled, gather_kv: true, gather_q: false, replace: null, polygraphy_class: AttentionLayerHint } ], dist_collectives: { group_size: 0, root: -1, nb_rank: 2, reduce_op: max, groups: [], polygraphy_class: DistCollective }, inputs: [ { name: input, seq_len_idx: 0, rank: 3, polygraphy_class: ShardTensor } ], outputs: [ { name: output, seq_len_idx: 0, rank: 3, polygraphy_class: ShardTensor } ], k_seq_len_idx: 3, v_seq_len_idx: 2, kv_rank: 4, polygraphy_class: ShardHints }顶层字段字段值含义parallelismCP并行类型CP/DP/PP。目前 Polygraphy 仅支持 CPattention_layers数组注意力层配置列表每个元素描述一个注意力层如何切分dist_collectives对象集合通信相关配置见下文inputs数组需要被 ReduceScatter 的张量列表outputs数组需要被 AllGather 的张量列表k_seq_len_idx3所有 K 张量中序列长度所在的维度索引非 0 时工具会在 DistCollective 节点前后自动插入 Transpose 把序列维度换到第一维v_seq_len_idx2所有 V 张量中序列长度所在的维度索引语义同上kv_rank4所有 K/V 张量的秩rank当模型无法推断维度且k_seq_len_idx/v_seq_len_idx非 0 时作为兜底polygraphy_classShardHints反序列化标记对应源码中 ShardHints 类attention_layers 字段每个注意力层配置包含q进入 QK^T 矩阵乘的 Q 张量名称本示例为q_scaled它在 create_onnx.py 中被显式命名作为 sharding 的锚点gather_kv是否在注意力之前 AllGather K/V。true表示每个 GPU 都持有完整的 K/V 以计算全局注意力gather_q是否 AllGather Q。false表示 Q 保持本地每个 GPU 只处理自己那一份分片polygraphy_class对应源码中的 AttentionLayerHint 类。注意仓库中的hint.json中额外出现replace: null和dist_collectives对象这是工具对 hint 序列化/反序列化往返后补充的默认字段。Polygraphy 工具文档multi_device/README.md给出的规范形式为attention_layers、inputs、outputs、k_seq_len_idx、v_seq_len_idx、kv_rank、reduce_scatter_reduce_op等顶层字段两种写法对 sharding 结果等效。dist_collectives 与归约算子本示例的hint.json中dist_collectives包含nb_rank模型将被切分到的 GPU 数量本示例为 2即 2 卡并行reduce_opReduceScatter 使用的归约操作本示例为max。对于分布输入的场景reduce 操作通常不影响语义正确性但必须在工具文档中登记的合法操作符内选择。在 Polygraphy 官方工具的规范写法中对应的顶层字段是group_size参与并行的 GPU 数0 表示使用全部可用 GPU和reduce_scatter_reduce_opreduce-scatter 节点的归约算子。inputs / outputs 字段每个张量配置包含nameONNX 模型中的张量名输入张量名input、输出张量名output均由 create_onnx.py 与 L225-L229 定义seq_len_idx序列长度所在的维度索引。本示例张量形状为(seq, batch, 4096)序列维度在索引 0因此为0若为非 0 索引工具会插入 Transpose 节点把序列维度调整到第一维rank张量的维度数量本示例为 3。当seq_len_idx非 0 且无法从模型推断出维度时作为兜底使用polygraphy_class对应源码中的 ShardTensor 类。Polygraphy 会根据这些提示为inputs中的张量插入 ReduceScatter为outputs中的张量插入 AllGather从而在单设备图上生成可在多卡上运行的多设备图。环境准备与依赖安装运行本示例需要满足以下前置条件一台配备多张 GPU的机器polygraphy 0.49.25提供multi-device shard子命令已安装 CUDA 与 TensorRT 的 Python 绑定tensorrt、cuda-bindings。示例目录下的 requirements.txt 列出了完整依赖numpy1.26.4 cuda-bindings mpi4py4.1.1 nccl4py0.1.1 onnx_graphsurgeon polygraphy0.49.25 torch其中mpi4py负责多进程编排与引擎/输入的广播nccl4py提供 NCCL 通信器用于 GPU 间集合通信onnx_graphsurgeonGraphSurgeon与polygraphy分别用于构建 ONNX 图和执行 sharding。安装命令pip3 install -r requirements.txt第一步生成单设备 ONNX 模型使用 GraphSurgeon Layer API 构建模型create_onnx.py使用 ONNX GraphSurgeon 的layer API从零构建自注意力模型而不是导出已有框架模型。脚本通过gs.Graph.register()把matmul、transpose、reshape、softmax、cast、sqrt、add、mul、div、pow、reduce_mean、shape_op、gather、unsqueeze、concat等 ONNX 算子注册为Graph的方法见 create_onnx.py使构图代码可以像搭积木一样链式调用。模型的核心结构build_attention_graphQ/K/V 线性投影输入(seq, batch, 4096)与三组4096x4096的 fp16 权重分别做 MatMul动态 Reshape 为多头布局(seq, batch, 4096) - (seq, batch, 32, 128)其中 32 个注意力头、每头 128 维NUM_HEADS 32、HEAD_DIM 128、HIDDEN_DIM 4096该 reshape 通过Shape - Gather - Unsqueeze - Concat动态计算目标形状以支持动态序列长度RMSNorm作用于 Q 和 K先 Cast 到 fp32计算x * rsqrt(mean(x^2) eps)其中eps 1e-6再乘上随机权重并转回 fp16缩放点积注意力Q 与 K 各自乘以sqrt(sqrt(1/HEAD_DIM))的缩放因子缩放被拆到 Q/K 两侧对应张量被命名为q_scaled随后执行QK^T - Softmax - Attn*V输出投影将注意力输出 reshape 回(seq, batch, 4096)后与输出权重 MatMul。生成命令python3 create_onnx.py --output attention_sd.onnx生成的模型具有以下特征输入/输出(sequence_length, batch_size, 4096)float1632 个注意力头每头 128 维包含 Q/K/V 投影、RMSNorm、缩放点积注意力与输出投影q_scaled与output张量被显式命名作为后续 sharding 的锚点。运行后脚本会打印节点数、初始器数量和 ONNX opset 版本模型 opset 为 17ir_version设为 8。第二步Sharding 生成多设备模型polygraphy multi-device shard attention_sd.onnx -s hint.json -o attention_md.onnx该命令把单设备模型attention_sd.onnx转换为多设备模型attention_md.onnx。处理过程的核心逻辑位于 shard.py工具解析 hint 文件遍历图中节点当发现DistCollective类型的节点时直接保留或重建源码中if node.op_type DistCollective的判断同时在需要的位置以gs.Node(opDistCollective, ...)创建新节点。最终attention_md.onnx中会包含 ReduceScatter 与 AllGather 算子设计用于在 2 张 GPU 上运行。对于 sharding 工具的完整文档与 hint 格式说明可阅读仓库内 Polygraphy multi-device 工具文档其中还提供了 hints 各字段parallelism、group_size、root、groups、attention_layers、inputs、outputs、k_seq_len_idx、v_seq_len_idx、kv_rank、reduce_scatter_reduce_op的逐项释义。需要特别说明的是polygraphy multi-device工具目前仅支持 CP 并行见 multi_device.py 工具定义 与 工具文档。运行示例单 GPU 基线运行单设备推理作为性能基线使用 sharding 前的attention_sd.onnxpython3 attention_mdtrt.py \ --onnx-path attention_sd.onnx \ --sequence-length 56320 \ --batch-size 1 \ --num-iterations 50多 GPU2 卡运行多设备推理通过mpirun启动 2 个进程每个进程绑定一张 GPU使用 sharding 后的attention_md.onnxmpirun -np 2 python3 attention_mdtrt.py \ --onnx-path attention_md.onnx \ --sequence-length 56320 \ --batch-size 1 \ --num-iterations 50使用指定的 libnccl.so如需强制加载特定版本的 NCCL 动态库例如系统默认 NCCL 与 nccl4py 不兼容时可通过LD_PRELOAD预加载LD_PRELOAD/path/to/libnccl.so mpirun -np 2 python3 attention_mdtrt.py \ --onnx-path attention_md.onnx推理程序运行机制深度剖析attention_mdtrt.py 是示例的核心推理脚本其运行机制可以从以下几个层面理解。1. 单进程/多进程的统一入口main()中如果环境中存在mpi4py则通过MPI.COMM_WORLD获取num_ranks进程/GPU 数与rank当前进程编号否则回退为单进程模式num_ranks 1、rank 0。脚本自动根据num_ranks决定执行单设备推理还是多设备推理因此同一个脚本文件同时承载单卡基线与多卡并行两条路径。2. 随机输入生成与广播输入数据只在 root rankrank 0上生成torch.rand((sequence_length, batch_size, 4096))转 float16 后转为连续 numpy 数组固定随机种子 42 以保证可复现。多设备模式下root 通过mpi_comm.bcast(input_data, rootroot)把完整输入广播给所有 rank——每个 rank 随后在本地持有的都是完整序列数据真正的序列切分由 ONNX 图中内嵌的 ReduceScatter 算子在设备端完成。3. NCCL 通信器初始化与 PyCapsule 传递多设备分支中的AttentionMD.setup_multidevice()attention_mdtrt.py完成以下工作通过cudart.cudaSetDevice(rank)将当前进程绑定到编号为rank的 GPU仅由 root rank 调用nccl.get_unique_id()生成 NCCL 通信 ID通过mpi_comm.bcast将通信 ID 广播给所有 rank每个 rank 调用nccl.Communicator.init(nranksnum_ranks, rankrank, unique_idnccl_comm_id)初始化参与集体通信的 NCCL 通信器。由于 TensorRT 的set_communicator接口接收的是 PyCapsule 而非 Python 对象脚本提供了communicator_to_capsule()辅助函数attention_mdtrt.py它校验nccl.core.Communicator对象仍存活ptr ! 0然后通过ctypes.pythonapi.PyCapsule_New将底层ncclComm_t句柄包装为名为ncclComm_t的 PyCapsule。4. 引擎构建与多进程广播AttentionMD.setup()中只有 root rank 调用build_serialized_network()构建并序列化引擎随后通过mpi_comm.bcast(engine_bin, rootroot)把序列化引擎广播到所有 rank各 rank 用trt.Runtime.deserialize_cuda_engine反序列化。这种单点构建、广播分发的策略避免了每个进程重复执行耗时的引擎构建。5. 通信器注入与动态形状反序列化后脚本通过context.set_communicator(capsule)把 NCCL 通信器注入执行上下文——这是 TensorRT 在推理时执行内嵌 DistCollective 算子的关键。随后对动态输入调用context.set_input_shape(input_name, actual_input_shape)指定实际形状通过allocate_buffers()attention_mdtrt.py分配主机与设备缓冲支持 BF16/HALF/FLOAT 等多种数据类型通过context.set_tensor_address绑定所有 I/O 张量地址。6. 性能计时与输出处理infer()中先执行一次 warmup再循环num_iterations次调用context.execute_async_v3并在每个流上同步最后输出平均耗时毫秒[Rank 0] Time spent in TRT attention: X.XX ms输出张量根据数据类型分别解析为 fp16 / bf16 / fp32 并 reshape 回(sequence_length, batch_size, 4096)。多卡模式下每个 rank 的输出形状仍是完整的序列长度——AllGather 算子已经在设备端重建了完整输出。可通过--save-output参数把输出保存为.npy文件仅 root rank 生效见 attention_mdtrt.py。7. 命令行参数一览参数默认值说明--onnx-path必填ONNX 模型路径单卡用attention_sd.onnx多卡用attention_md.onnx--sequence-length56320输入序列长度--batch-size1批大小--num-iterations50计时推理迭代次数--save-outputNone将输出张量保存为.npy文件仅 root rank单卡与多卡的正确使用方式由于模型构建时输入优化 profile 的形状范围为(1, 1, 4096) - (56320, 1, 4096) - (56320, 1, 4096)见 attention_mdtrt.py示例按默认序列长度 56320 运行即可落在优化形状上。实际使用中请注意单卡基线必须使用 sharding 前的attention_sd.onnx多卡推理必须使用 sharding 后的attention_md.onnx且进程数-np应与 hint 中nb_rank本示例为 2一致多卡模式强依赖mpi4py与nccl4py缺失时脚本会在启动时明确报错退出构建引擎时设置了 16GB workspace 与 1GB tactic shared memory 上限attention_mdtrt.py请确保 GPU 显存充足序列长度可调整但需同时保证模型输入在优化 profile 允许的动态范围内。相关资源示例目录samples/python/attention_mdtrt含 README.md、attention_mdtrt.py、create_onnx.py、hint.json、requirements.txtPolygraphy sharding 工具文档tools/Polygraphy/polygraphy/tools/multi_device/README.mdsharding 工具实现shard.py 与 hint 数据结构定义 multi_device.py许可证条款见仓库根目录 LICENSE 与 NOTICEChangelog 与已知问题该示例的演进记录见 README.md2026 年 4 月新增create_onnx.py改用 GraphSurgeon layer API 生成 ONNX 模型新增--save-output标志用于保存推理输出文档补充 DistCollective 与 sharding 机制说明2026 年 1 月示例首次发布。目前该示例没有已知问题Known Issues: None。【免费下载链接】TensorRTNVIDIA® TensorRT™ is an SDK for high-performance deep learning inference on NVIDIA GPUs. This repository contains the open source components of TensorRT.项目地址: https://gitcode.com/GitHub_Trending/tens/TensorRT创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考