diff --git a/tests/compile/passes/distributed/test_fusion_all_reduce.py b/tests/compile/passes/distributed/test_fusion_all_reduce.py index 17c4e30959b..b86018a7555 100644 --- a/tests/compile/passes/distributed/test_fusion_all_reduce.py +++ b/tests/compile/passes/distributed/test_fusion_all_reduce.py @@ -13,6 +13,7 @@ from vllm._custom_ops import cutlass_scaled_fp4_mm, scaled_fp4_quant from vllm.compilation.passes.fusion.allreduce_rms_fusion import ( AllReduceFusionPass, RocmAiterAllReduceFusionPass, + _select_flashinfer_allreduce_use_oneshot, ) from vllm.compilation.passes.fx_utils import find_op_nodes from vllm.compilation.passes.utility.fix_functionalization import ( @@ -48,6 +49,35 @@ from vllm.utils.torch_utils import set_random_seed DEVICE_TYPE = current_platform.device_type +@pytest.mark.parametrize( + ("workspace_backend", "device_capability", "world_size", "tensor_size", "expected"), + [ + ("mnnvl", 103, 8, 2 * 1024 * 1024, None), + ("trtllm", 103, 8, 2 * 1024 * 1024, True), + ("trtllm", 103, 8, 2 * 1024 * 1024 + 1, False), + ("trtllm", 100, 4, 4 * 1024 * 1024, True), + ("trtllm", 100, 4, 4 * 1024 * 1024 + 1, False), + ("trtllm", None, 8, 128 * 1024 * 1024, True), + ], +) +def test_select_flashinfer_allreduce_use_oneshot( + workspace_backend: str, + device_capability: int | None, + world_size: int, + tensor_size: int, + expected: bool | None, +): + assert ( + _select_flashinfer_allreduce_use_oneshot( + workspace_backend, + device_capability, + world_size, + tensor_size, + ) + is expected + ) + + class TestAllReduceRMSNormModel(torch.nn.Module): def __init__( self, diff --git a/vllm/compilation/passes/fusion/allreduce_rms_fusion.py b/vllm/compilation/passes/fusion/allreduce_rms_fusion.py index f4962424d7f..ab400028925 100644 --- a/vllm/compilation/passes/fusion/allreduce_rms_fusion.py +++ b/vllm/compilation/passes/fusion/allreduce_rms_fusion.py @@ -128,6 +128,29 @@ _FI_ALLREDUCE_ONE_SHOT_MAX_SIZES_MB: dict[int, dict[int, float]] = { }, } +MiB = 1024 * 1024 + + +def _select_flashinfer_allreduce_use_oneshot( + workspace_backend: str, + device_capability: int | None, + world_size: int, + current_tensor_size: int, +) -> bool | None: + if workspace_backend == "mnnvl": + # FlashInfer sizes MNNVL workspaces around its own AUTO strategy. + # Forcing vLLM's per-rank threshold can request one-shot for tensors + # larger than the MNNVL one-shot workspace. + return None + + if device_capability is None: + max_one_shot_size = None + else: + max_one_shot_size = _FI_ALLREDUCE_ONE_SHOT_MAX_SIZES_MB.get( + device_capability, {} + ).get(world_size) + return max_one_shot_size is None or current_tensor_size <= max_one_shot_size * MiB + if flashinfer_comm is not None: from vllm.distributed.device_communicators.flashinfer_all_reduce import ( @@ -138,8 +161,6 @@ if flashinfer_comm is not None: ar_fusion_patterns = flashinfer_comm.AllReduceFusionPattern - MiB = 1024 * 1024 - def call_trtllm_fused_allreduce_norm( allreduce_in: torch.Tensor, residual: torch.Tensor, @@ -174,16 +195,6 @@ if flashinfer_comm is not None: ) curr_device = current_platform.get_device_capability() device_capability = curr_device.to_int() if curr_device is not None else None - # Get one shot input size limit for the current world size - # for the current device capability - max_one_shot_size = _FI_ALLREDUCE_ONE_SHOT_MAX_SIZES_MB.get( - device_capability, # type: ignore[arg-type, unused-ignore] - {}, - ).get(world_size, None) - # Use one shot if no max size is specified - use_oneshot = ( - max_one_shot_size is None or current_tensor_size <= max_one_shot_size * MiB - ) # Select workspace based on pattern: quant patterns use the # trtllm quant workspace, non-quant patterns use the primary workspace. @@ -205,6 +216,12 @@ if flashinfer_comm is not None: assert workspace is not None, ( "Flashinfer allreduce workspace must be initialized when using flashinfer" ) + use_oneshot = _select_flashinfer_allreduce_use_oneshot( + workspace.backend, + device_capability, + world_size, + current_tensor_size, + ) assert flashinfer_comm is not None if norm_out is None: norm_out = allreduce_in @@ -248,7 +265,7 @@ if flashinfer_comm is not None: # the end for the one-shot path; the two-shot path is synchronized # and keeps the early completion. Related one-shot instability in # the same kernel: flashinfer-ai/flashinfer#1223. - trigger_completion_at_end=use_oneshot + trigger_completion_at_end=(use_oneshot is True) or num_tokens > PDL_ADVANCE_LAUNCH_TOKENS, )