forked from Karylab-cklius/vllm
[Bugfix][Distributed] Delegate MNNVL allreduce one-shot selection (#47589)
Signed-off-by: jesco-absolut <team@srswti.com> Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent
095adf1fdc
commit
598d51153a
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user