flashinfer.comm.trtllm_create_ipc_workspace_for_all_reduce_fusion¶
- flashinfer.comm.trtllm_create_ipc_workspace_for_all_reduce_fusion(tp_rank: int, tp_size: int, max_token_num: int, hidden_dim, use_fp32_lamport: bool = False, group: ProcessGroup | None = None, create_metadata: bool = False, comm_backend: CommBackend | None = None, use_symm_dev_mem: bool = False) Tuple[List[List[int]], Tensor] | Tuple[List[List[int]], Tensor, dict] | Tuple[List[List[int]], Tensor, List[SymmDeviceMemory], dict]¶
参数: - tp_rank: 当前进程的 rank。 - tp_size: 进程组的大小。 - max_token_num: 序列中的最大 token 数量。 - hidden_dim: 隐藏状态的维度。 - use_fp32_lamport: 如果为 True,则在 allreduce fusion 中使用 fp32 数据类型。 - group: 要使用的进程组。 - create_metadata: 如果为 True,则将 metadata dict 作为第三个元素返回 (默认: False)。 - comm_backend: 要使用的通信后端。 - use_symm_dev_mem: 如果为 True,则为工作区使用对称设备内存。
返回值: - 如果 create_metadata=False: (ipc_handles, workspace_tensor) - 如果 create_metadata=True: 且 use_symm_dev_mem=False: (ipc_handles, workspace_tensor, metadata)
其中 metadata 包含: tp_rank, tp_size, max_token_num, hidden_dim, use_fp32_lamport, buffer_size, flag_size, lamport_comm_size, lamport_buffer_size
如果 create_metadata=True: 且 use_symm_dev_mem=True: (ipc_handles, workspace_tensor, mem_handles,metadata) 其中 metadata 包含: tp_rank, tp_size, max_token_num, hidden_dim, use_fp32_lamport, buffer_size, flag_size, lamport_comm_size, lamport_buffer_size 并且 mem_handles 是 SymmDeviceMemory 对象的列表。
注意: 可选参数目前使 API 显得笨拙。 这将在未来进行重构,以牺牲向后兼容性为代价,其中默认行为将是 create_metadata=True 且 use_symm_dev_mem=True。
注意: 我们将为 trtllm_custom_all_reduce_fusion 初始化 3 个 IPC 缓冲区。 它们的大小如下: [buffer_size, flag_size, lamport_buffer_size * 3] 其中: - buffer_size: tp_size * max_token_num * hidden_dim * sizeof(half) - flag_size: tp_size * BarrierFlagCount * sizeof(int) - lamport_buffer_size: tp_size * max_token_num * tp_size * hidden_dim * sizeof(half)
其中 sizeof(elem) = 2 (fp16/bf16) 或 4 (use_fp32_lamport=True 时的 fp32)
工作区作为 AllReduceFusionParams 中的 workspace 字段传递。
我们在这里互换使用 tp_size 和 world_size (allReduceFusion)。
参考: trtllm, cpp/tensorrt_llm/kernels/communicationKernels/allReduceWorkspace.cu, Workspace init