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,则需要)