flashinfer.comm.trtllm_allreduce_fusion

flashinfer.comm.trtllm_allreduce_fusion(allreduce_in: Tensor, world_size: int, world_rank: int, token_num: int, hidden_dim: int, workspace_ptrs: Tensor, launch_with_pdl: bool, trigger_completion_at_end: bool, fp32_acc: bool, pattern_code: AllReduceFusionPattern, use_oneshot: bool | None, allreduce_out: Tensor | None, residual_in: Tensor | None, residual_out: Tensor | None, norm_out: Tensor | None, quant_out: Tensor | None, scale_out: Tensor | None, rms_gamma: Tensor | None, rms_eps: float | None, scale_factor: Tensor | float | None, layout_code: QuantizationSFLayout | None, metadata: dict | None = None) None

参数: - allreduce_in: 输入张量。 [token_num, hidden_dim] - world_size: 进程组的大小。 - world_rank: 当前进程的等级。 - token_num: 序列中的 token 数量。 - hidden_dim: 隐藏状态的维度。 - workspace_ptrs: 工作区指针。 - launch_with_pdl: 是否使用 pdl 启动。 - use_oneshot: 是否使用 oneshot。 如果为 None,将使用内部启发式方法。 - trigger_completion_at_end: 是否在结束时触发完成。 - fp32_acc: 是否使用 fp32 累积。 - pattern_code: 模式代码。 - allreduce_out: 输出张量。 [token_num, hidden_dim] - residual_in: 残差输入张量。 [token_num, hidden_dim] - residual_out: 残差输出张量。 [token_num, hidden_dim] - norm_out: 归一化输出张量。 [token_num, hidden_dim] - quant_out: 量化输出张量。 [token_num, hidden_dim] - scale_out: 缩放输出张量。 初始化参考: tests/comm/test_trtllm_allreduce_fusion.py - rms_gamma: rms gamma 张量。 [hidden_dim] - rms_eps: rms epsilon 值。 - scale_factor: 缩放因子。 为了 cudaGraphs 的安全,它应该是一个张量。 - layout_code: 布局代码。 - metadata: create_ipc_workspace_for_all_reduce_fusion 返回的可选工作区元数据字典。

如果提供,则验证 token_num <= max_token_num、world_size == tp_size 和 hidden_dim == workspace hidden_dim。 如果验证失败,则引发 ValueError。