[SRT] fix flashInfer allreduce fusion not used on blackwell (#26197)

This commit is contained in:
Mick
2026-05-25 18:07:15 +08:00
committed by GitHub
parent a4db563c87
commit c05756da7a
@@ -1,4 +1,5 @@
import contextlib
import inspect
import logging
import platform
from typing import Optional, Tuple
@@ -30,6 +31,9 @@ logger = logging.getLogger(__name__)
_flashinfer_comm = None
_TorchDistBackend = None
_flashinfer_allreduce_unavailable = False
_flashinfer_create_workspace_supports_group = False
_flashinfer_create_workspace_supports_comm_backend = False
_flashinfer_allreduce_supports_trigger_completion = False
_posix_transport_override_logged = False
@@ -106,6 +110,17 @@ if is_flashinfer_available():
comm, "create_allreduce_fusion_workspace"
):
_flashinfer_comm = comm
workspace_params = inspect.signature(
comm.create_allreduce_fusion_workspace
).parameters
allreduce_params = inspect.signature(comm.allreduce_fusion).parameters
_flashinfer_create_workspace_supports_group = "group" in workspace_params
_flashinfer_create_workspace_supports_comm_backend = (
"comm_backend" in workspace_params
)
_flashinfer_allreduce_supports_trigger_completion = (
"trigger_completion_at_end" in allreduce_params
)
else:
_flashinfer_allreduce_unavailable = True
logger.warning(
@@ -383,14 +398,15 @@ class FlashInferWorkspaceManager:
hidden_dim=hidden_dim,
dtype=dtype,
force_oneshot_support=bool(use_oneshot),
# Pin the symmetric-memory rendezvous to the actual
# subgroup. Without this, flashinfer >=0.6.10 falls back
# to WORLD and TP/EP/CP subgroup peers get addressed
# incorrectly (kernel hangs in cuda-graph warmup).
group=device_group,
)
create_workspace = _flashinfer_comm.create_allreduce_fusion_workspace
if _flashinfer_create_workspace_supports_group:
# Pin the symmetric-memory rendezvous to the actual subgroup.
# Older FlashInfer releases only support comm_backend.
kwargs["group"] = device_group
if (
_TorchDistBackend is not None
and _flashinfer_create_workspace_supports_comm_backend
and device_group is not None
and cpu_group is not None
):
@@ -398,9 +414,7 @@ class FlashInferWorkspaceManager:
device_group=device_group, cpu_group=cpu_group
)
with _flashinfer_posix_fd_transport_override_if_needed():
self.workspace = _flashinfer_comm.create_allreduce_fusion_workspace(
**kwargs
)
self.workspace = create_workspace(**kwargs)
except Exception as e:
_flashinfer_allreduce_unavailable = True
logger.warning(
@@ -669,12 +683,11 @@ def flashinfer_allreduce_residual_rmsnorm(
norm_out = torch.empty_like(input_tensor)
workspace_manager = _get_workspace_manager(use_attn_tp_group)
_flashinfer_comm.allreduce_fusion(
kwargs = dict(
input=input_tensor,
workspace=workspace_manager.workspace,
pattern=_flashinfer_comm.AllReduceFusionPattern.kARResidualRMSNorm,
launch_with_pdl=True,
trigger_completion_at_end=trigger_completion_at_end,
residual_out=residual_out,
norm_out=norm_out,
residual_in=residual,
@@ -683,6 +696,9 @@ def flashinfer_allreduce_residual_rmsnorm(
use_oneshot=use_oneshot,
fp32_acc=fp32_acc,
)
if _flashinfer_allreduce_supports_trigger_completion:
kwargs["trigger_completion_at_end"] = trigger_completion_at_end
_flashinfer_comm.allreduce_fusion(**kwargs)
return norm_out, residual_out