flashinfer.comm.trtllm_custom_all_reduce¶
- flashinfer.comm.trtllm_custom_all_reduce(inp: Tensor, out: Tensor, tp_size: int, tp_rank: int, token_num: int, fusion_op_code: AllReduceFusionOp, strategy_code: AllReduceStrategyType, config_code: AllReduceStrategyConfig, launch_with_pdl: bool, flag_value: int, peer_comm_buffer_ptrs: Tensor, peer_barrier_ptrs_in: Tensor, peer_barrier_ptrs_out: Tensor, bias: Tensor | None, residual: Tensor | None, weight: Tensor | None, weight_pre_residual_norm: Tensor | None, eps: float | None, intermediate_buffer: Tensor | None, lamport_peer_comm_buffer_ptrs_0: Tensor | None, lamport_peer_comm_buffer_ptrs_1: Tensor | None, lamport_peer_comm_buffer_ptrs_2: Tensor | None) None¶
参数: - inp: 输入张量。 [token_num, hidden_dim] - out: 输出张量。 [token_num, hidden_dim] - tp_size: 进程组的大小。 - tp_rank: 当前进程的等级。 - token_num: 序列中的token数量。 - fusion_op_code: 融合操作码。 - strategy_code: 策略码。 - config_code: 配置码。 - launch_with_pdl: 是否使用pdl启动。 - flag_value: 标志值。 - peer_comm_buffer_ptrs: 对等通信缓冲区指针。 - peer_barrier_ptrs_in: 对等屏障指针输入。 - peer_barrier_ptrs_out: 对等屏障指针输出。 - bias: 偏置张量。 [hidden_dim] - residual: 残差张量。 [token_num, hidden_dim] - weight: 权重张量。 [hidden_dim] - weight_pre_residual_norm: 残差归一化前的权重张量。 [hidden_dim] - eps: epsilon值。 - intermediate_buffer: 中间缓冲区张量。 - lamport_peer_comm_buffer_ptrs_0: lamport对等通信缓冲区指针0。 - lamport_peer_comm_buffer_ptrs_1: lamport对等通信缓冲区指针1。 - lamport_peer_comm_buffer_ptrs_2: lamport对等通信缓冲区指针2。