flashinfer.comm.trtllm_mnnvl_ar.trtllm_mnnvl_all_reduce¶
- flashinfer.comm.trtllm_mnnvl_ar.trtllm_mnnvl_all_reduce(inp: Tensor, multicast_buffer_ptr: int, buffer_ptrs_dev: int, buffer_M: int, buffer_flags_mnnvl: Tensor, nranks: int, rank: int, wait_for_results: bool, launch_with_pdl: bool, out: Tensor | None = None) None¶
在多个 GPU 上执行多节点 NVLink 全归约操作。
此函数使用 NVIDIA 的多节点 NVLink (MNNVL) 技术执行全归约(求和)操作,以有效地组合多个 GPU 和节点上的张量。
有 3 个步骤:1. 将每个 GPU 的输入分片散播到正确的单播缓冲区 2. 在每个 GPU 上执行全归约 3. 将结果广播到所有 GPU
- 参数:
inp – 本地输入分片
multicast_buffer_ptr – 指向多播缓冲区的整数指针
buffer_ptrs_dev – 指向设备缓冲区指针的整数
buffer_M – 最大元素数 // hidden_dim
buffer_flags_mnnvl – 包含缓冲区状态标志的张量
nranks – 参与全归约的总 rank 数
rank – 当前进程 rank
wait_for_results – 如果为 True,则将结果存储到 out
launch_with_pdl – 如果为 True,则使用程序化依赖启动
out ([可选]) – 存储结果的输出张量(如果 wait_for_results 为 True,则需要)