[refactor] Move MLP collective flags onto ForwardFlags (#30802)
This commit is contained in:
@@ -696,6 +696,12 @@ class TestForwardFlags(_IsolatedServerArgs):
|
||||
x = x + 1
|
||||
if fwd.is_extend_in_batch:
|
||||
x = x + 2
|
||||
if fwd.fuse_mlp_allreduce:
|
||||
x = x + 4
|
||||
if fwd.mlp_reduce_scatter:
|
||||
x = x + 8
|
||||
if fwd.flashinfer_trtllm_bypass:
|
||||
x = x + 16
|
||||
return x
|
||||
|
||||
self.assertEqual(probe(torch.zeros(())).item(), 0)
|
||||
@@ -704,6 +710,13 @@ class TestForwardFlags(_IsolatedServerArgs):
|
||||
get_forward().set("is_extend_in_batch", True)
|
||||
self.assertEqual(probe(torch.zeros(())).item(), 2)
|
||||
get_forward().set("is_extend_in_batch", False)
|
||||
with get_forward().scoped(
|
||||
fuse_mlp_allreduce=True,
|
||||
mlp_reduce_scatter=True,
|
||||
flashinfer_trtllm_bypass=True,
|
||||
):
|
||||
self.assertEqual(probe(torch.zeros(())).item(), 28)
|
||||
self.assertEqual(probe(torch.zeros(())).item(), 0)
|
||||
|
||||
def test_graph_visible_flags_are_process_visible_across_threads(self):
|
||||
# Documented divergence from the contextvar-backed flags: plain slots
|
||||
@@ -812,6 +825,38 @@ class TestForwardFlags(_IsolatedServerArgs):
|
||||
self.assertIs(get_forward().moe_output_buffer, sentinel)
|
||||
self.assertIsNone(get_forward().moe_output_buffer)
|
||||
|
||||
def test_mlp_comm_forward_flags(self):
|
||||
"""Decoder-published MLP collective flags: scoped restore + skip helpers."""
|
||||
from sglang.srt.layers.moe.utils import (
|
||||
should_skip_mlp_all_reduce,
|
||||
should_skip_post_experts_all_reduce,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_forward
|
||||
|
||||
reset_context()
|
||||
fwd = get_forward()
|
||||
self.assertFalse(fwd.fuse_mlp_allreduce)
|
||||
self.assertFalse(fwd.mlp_reduce_scatter)
|
||||
self.assertFalse(fwd.flashinfer_trtllm_bypass)
|
||||
self.assertFalse(should_skip_mlp_all_reduce())
|
||||
|
||||
with fwd.scoped(fuse_mlp_allreduce=True):
|
||||
self.assertTrue(fwd.fuse_mlp_allreduce)
|
||||
self.assertTrue(should_skip_mlp_all_reduce())
|
||||
# Fusion alone is enough to skip post-experts AR.
|
||||
self.assertTrue(should_skip_post_experts_all_reduce(is_tp_path=True))
|
||||
self.assertFalse(fwd.fuse_mlp_allreduce)
|
||||
self.assertFalse(should_skip_mlp_all_reduce())
|
||||
|
||||
with fwd.scoped(mlp_reduce_scatter=True):
|
||||
self.assertTrue(fwd.mlp_reduce_scatter)
|
||||
self.assertTrue(should_skip_mlp_all_reduce())
|
||||
self.assertFalse(fwd.mlp_reduce_scatter)
|
||||
|
||||
with fwd.scoped(flashinfer_trtllm_bypass=True):
|
||||
self.assertTrue(fwd.flashinfer_trtllm_bypass)
|
||||
self.assertFalse(fwd.flashinfer_trtllm_bypass)
|
||||
|
||||
|
||||
class TestPublishLifecycle(_IsolatedServerArgs):
|
||||
"""Publish installs the resolved server_args and seeds the capture tier."""
|
||||
|
||||
Reference in New Issue
Block a user