flashinfer.comm.trtllm_mnnvl_ar.trtllm_mnnvl_fused_allreduce_rmsnorm¶
- flashinfer.comm.trtllm_mnnvl_ar.trtllm_mnnvl_fused_allreduce_rmsnorm(prenorm_output: Tensor, normed_output: Tensor, shard_input: Tensor, multicast_buffer_ptr: int, buffer_ptrs_dev: int, unicast_ptr: int, buffer_M: int, buffer_flags_mnnvl: Tensor, nranks: int, rank: int, gamma: Tensor, epsilon: float, residual: Tensor, launch_with_pdl: bool) None¶
执行 MNNVL 两阶段 Allreduce + RMSNorm。
此函数通过首先调用 trtllm_mnnvl_all_reduce 在 shard_input 上执行多节点归约(求和)操作。 之后,它从多播缓冲区直接读取归约结果,执行 RMSNorm。 注意:对于当前 rank,多播缓冲区与单播缓冲区相同。
- 参数:
prenorm_output – prenorm 结果的输出张量
normed_output – 归一化结果的输出张量
shard_input – 输入张量分片
multicast_buffer_ptr – 多播缓冲区的整数指针地址
buffer_ptrs_dev – 设备缓冲区指针的整数指针地址
unicast_ptr – 单播缓冲区的整数指针地址
buffer_M – 最大元素数量 // hidden_dim
buffer_flags_mnnvl – 同步的缓冲区标志
nranks – 张量并行组中的 rank 数量
rank – 张量并行组中的当前 rank
gamma – RMSNorm 的 gamma(归一化权重)参数
epsilon – RMSNorm 的 epsilon 参数
residual – 要添加的残差张量
launch_with_pdl – 是否使用 PDL 启动