flashinfer.comm.trtllm_create_ipc_workspace_for_all_reduce

flashinfer.comm.trtllm_create_ipc_workspace_for_all_reduce(rank: int, tp_size: int, max_token_num: int, hidden_dim, group: ProcessGroup | None = None) List[List[int]]

参数: - rank: 当前进程的 rank。 - tp_size: 进程组的大小。 - max_token_num: 序列中的最大 token 数量。 - hidden_dim: 隐藏状态的维度。 - group: 要使用的进程组。

注意: 此函数用于创建 all reduce 的工作区。工作区是 IPC handle 的列表。在调用 trtllm_custom_all_reduce 之前应初始化工作区。在调用 trtllm_custom_all_reduce 之后应销毁工作区。在相同的配置下,工作区可以重用于多个 all reduce 调用。

我们将为 trtllm_custom_all_reduce 初始化 7 个 IPC 缓冲区。它们的大小如下:[buffer_size, buffer_size, flag_size, flag_size, lamport_buffer_size, lamport_buffer_size, lamport_buffer_size] 其中: - buffer_size: tp_size * max_token_num * hidden_dim * sizeof(float) * (maxBeamWidth) - flag_size: (MAX_ALL_REDUCE_BLOCKS + 1) * sizeof(uint32_t) * tp_size * 2 - lamport_buffer_size: tp_size * LamportTokenNumThreshold * tp_size * hidden_dim * sizeof(half)

它们用于:ipcHandles[0] - peer_comm_buffer_ptrs ipcHandles[2] - peer_barrier_ptrs_in ipcHandles[3] - peer_barrier_ptrs_out ipcHandles[4] - lamport_peer_comm_buffer_ptrs[0:tp_size] ipcHandles[5] - lamport_peer_comm_buffer_ptrs[tp_size:tp_size * 2] ipcHandles[6] - lamport_peer_comm_buffer_ptrs[tp_size * 2:tp_size * 3]

我们在此处互换使用 tp_size 和 world_size (customAllReduce)。

参考: trtllm, cpp/tests/unit_tests/kernels/allReduce/allReduceKernelTest.cu, Workspace init