vllm.models.common.ops.fused_allreduce_rms_norm
¶
Fused all-reduce + residual-add + RMSNorm for eager model paths.
This recovers a fusion that vLLM's torch.compile passes would normally do but that doesn't fire for models running eager (or under a breakable CUDA graph).
Functions:
-
fused_allreduce_rms_norm–All-reduce + add residual + (standard) RMSNorm, fused via flashinfer.
fused_allreduce_rms_norm(hidden_states, residual, norm)
¶
All-reduce + add residual + (standard) RMSNorm, fused via flashinfer.
hidden_states is the per-rank partial output of a row-parallel linear
run with reduce_results=False; norm is the RMSNorm applied right
after. Returns (normed_output, new_residual), equivalent to
norm(all_reduce(hidden_states), residual). Falls back to an explicit
all-reduce + RMSNorm when the flashinfer fast path is unavailable.