[bugfix][AMD] Disable aiter allreduce+RMSNorm fusion under DP attention / EP (#27835)
This commit is contained in:
@@ -796,6 +796,8 @@ class LayerCommunicator:
|
|||||||
_use_aiter
|
_use_aiter
|
||||||
and batch_size > 0
|
and batch_size > 0
|
||||||
and get_parallel().tp_size != 6
|
and get_parallel().tp_size != 6
|
||||||
|
and not is_dp_attention_enabled()
|
||||||
|
and get_moe_a2a_backend().is_none()
|
||||||
and get_global_server_args().enable_aiter_allreduce_fusion
|
and get_global_server_args().enable_aiter_allreduce_fusion
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -3,12 +3,18 @@ import os
|
|||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import tempfile
|
import tempfile
|
||||||
|
import types
|
||||||
import unittest
|
import unittest
|
||||||
|
from contextlib import ExitStack
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers import communicator as comm
|
||||||
|
from sglang.srt.layers.communicator import LayerCommunicator, ScatterMode
|
||||||
from sglang.test.ci.ci_register import register_amd_ci
|
from sglang.test.ci.ci_register import register_amd_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
register_amd_ci(est_time=240, suite="stage-c-test-large-8-gpu-amd")
|
register_amd_ci(est_time=240, suite="stage-c-test-large-8-gpu-amd")
|
||||||
|
|
||||||
@@ -342,6 +348,131 @@ class TestAiterAllreduceFusionAmd(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _fake_self(*, mlp_mode=ScatterMode.TP_ATTN_FULL, is_last_layer=False, tp_size=8):
|
||||||
|
"""Minimal stand-in for a LayerCommunicator with the fields the gate reads."""
|
||||||
|
return types.SimpleNamespace(
|
||||||
|
_speculative_algo=None,
|
||||||
|
layer_scatter_modes=types.SimpleNamespace(mlp_mode=mlp_mode),
|
||||||
|
is_last_layer=is_last_layer,
|
||||||
|
_context=types.SimpleNamespace(tp_size=tp_size),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _fake_forward_batch(batch_size=8):
|
||||||
|
return types.SimpleNamespace(input_ids=types.SimpleNamespace(shape=(batch_size,)))
|
||||||
|
|
||||||
|
|
||||||
|
class TestAiterAllreduceFusionGate(CustomTestCase):
|
||||||
|
"""Pure-logic coverage of the aiter all-reduce + RMSNorm fusion gate.
|
||||||
|
|
||||||
|
Covers ``LayerCommunicator.should_fuse_mlp_allreduce_with_next_layer``,
|
||||||
|
specifically the AMD/aiter branch guards that disable the fused path under
|
||||||
|
DP attention or an expert-parallel A2A backend (e.g. mori). Without those
|
||||||
|
guards the fused custom all-reduce is invoked during CUDA graph capture in
|
||||||
|
those configs and crashes in ``custom_all_reduce.flush_graph_buffers``.
|
||||||
|
|
||||||
|
The gate is pure decision logic, so the test stubs out the module-level
|
||||||
|
dependencies and invokes the method on a minimal fake instance. No GPU or
|
||||||
|
distributed initialization is required.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _evaluate_gate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
dp_attention,
|
||||||
|
a2a_is_none,
|
||||||
|
aiter_enabled=True,
|
||||||
|
use_aiter=True,
|
||||||
|
tp_world_size=8,
|
||||||
|
mlp_mode=ScatterMode.TP_ATTN_FULL,
|
||||||
|
is_last_layer=False,
|
||||||
|
tp_size=8,
|
||||||
|
):
|
||||||
|
"""Run the gate with the aiter branch isolated (flashinfer forced off)."""
|
||||||
|
server_args = types.SimpleNamespace(enable_aiter_allreduce_fusion=aiter_enabled)
|
||||||
|
a2a_backend = types.SimpleNamespace(is_none=lambda: a2a_is_none)
|
||||||
|
|
||||||
|
with ExitStack() as stack:
|
||||||
|
stack.enter_context(
|
||||||
|
mock.patch.object(comm, "is_enable_moe_cp_allgather", lambda: False)
|
||||||
|
)
|
||||||
|
stack.enter_context(
|
||||||
|
mock.patch.object(
|
||||||
|
comm,
|
||||||
|
"get_attn_tp_context",
|
||||||
|
lambda: types.SimpleNamespace(input_scattered=False),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# Force the NVIDIA/flashinfer term off so the aiter branch decides.
|
||||||
|
stack.enter_context(
|
||||||
|
mock.patch.object(
|
||||||
|
comm, "apply_flashinfer_allreduce_fusion", lambda batch_size: False
|
||||||
|
)
|
||||||
|
)
|
||||||
|
stack.enter_context(mock.patch.object(comm, "_use_aiter", use_aiter))
|
||||||
|
stack.enter_context(
|
||||||
|
mock.patch.object(
|
||||||
|
comm,
|
||||||
|
"get_parallel",
|
||||||
|
lambda: types.SimpleNamespace(tp_size=tp_world_size),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
stack.enter_context(
|
||||||
|
mock.patch.object(comm, "get_global_server_args", lambda: server_args)
|
||||||
|
)
|
||||||
|
stack.enter_context(
|
||||||
|
mock.patch.object(comm, "is_dp_attention_enabled", lambda: dp_attention)
|
||||||
|
)
|
||||||
|
stack.enter_context(
|
||||||
|
mock.patch.object(comm, "get_moe_a2a_backend", lambda: a2a_backend)
|
||||||
|
)
|
||||||
|
|
||||||
|
fake_self = _fake_self(
|
||||||
|
mlp_mode=mlp_mode, is_last_layer=is_last_layer, tp_size=tp_size
|
||||||
|
)
|
||||||
|
return LayerCommunicator.should_fuse_mlp_allreduce_with_next_layer(
|
||||||
|
fake_self, _fake_forward_batch()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_dense_tp_fuses(self):
|
||||||
|
# Baseline supported path: dense TP, no DP attention, no EP backend.
|
||||||
|
self.assertTrue(self._evaluate_gate(dp_attention=False, a2a_is_none=True))
|
||||||
|
|
||||||
|
def test_dp_attention_disables_fusion(self):
|
||||||
|
# The fix: DP attention has no dense TP all-reduce to fuse.
|
||||||
|
self.assertFalse(self._evaluate_gate(dp_attention=True, a2a_is_none=True))
|
||||||
|
|
||||||
|
def test_ep_backend_disables_fusion(self):
|
||||||
|
# The fix: with an EP A2A backend (e.g. mori) the reduction lives in
|
||||||
|
# combine(), not a TP all-reduce.
|
||||||
|
self.assertFalse(self._evaluate_gate(dp_attention=False, a2a_is_none=False))
|
||||||
|
|
||||||
|
def test_dp_attention_and_ep_disables_fusion(self):
|
||||||
|
# The crashing config from the TP8+EP8+mori repro.
|
||||||
|
self.assertFalse(self._evaluate_gate(dp_attention=False, a2a_is_none=False))
|
||||||
|
self.assertFalse(self._evaluate_gate(dp_attention=True, a2a_is_none=False))
|
||||||
|
|
||||||
|
def test_flag_off_disables_fusion(self):
|
||||||
|
# Sanity: the gate still respects the opt-in flag on the dense path.
|
||||||
|
self.assertFalse(
|
||||||
|
self._evaluate_gate(
|
||||||
|
dp_attention=False, a2a_is_none=True, aiter_enabled=False
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_last_layer_disables_fusion(self):
|
||||||
|
self.assertFalse(
|
||||||
|
self._evaluate_gate(
|
||||||
|
dp_attention=False, a2a_is_none=True, is_last_layer=True
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_tp1_disables_fusion(self):
|
||||||
|
self.assertFalse(
|
||||||
|
self._evaluate_gate(dp_attention=False, a2a_is_none=True, tp_size=1)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
if "--residual-accuracy" in sys.argv:
|
if "--residual-accuracy" in sys.argv:
|
||||||
_run_residual_accuracy_check()
|
_run_residual_accuracy_check()
|
||||||
|
|||||||
Reference in New Issue
Block a user