Give the attention-DP width and rank one home (#40067)
This commit is contained in:
@@ -65,7 +65,9 @@ _DP = "sglang.srt.layers.dp_attention"
|
||||
# anything, so there is nothing to derive them from. The quotients used to be
|
||||
# in this table and are not any more -- `attn_tp_size` and its siblings are
|
||||
# functions of the configured leaves, and `TestDerivedWidthsComeFromTheLeaves`
|
||||
# is what pins them.
|
||||
# is what pins them. `attn_dp_rank` is not here either: no group coordinator
|
||||
# knows it, so it is stamped when the attention topology is initialized and
|
||||
# `TestStampedRanks` is what pins it.
|
||||
SIZE_RANK_DELEGATIONS = [
|
||||
("world_size", f"{_PS}.get_world_size"),
|
||||
("world_rank", f"{_PS}.get_world_rank"),
|
||||
@@ -77,7 +79,6 @@ SIZE_RANK_DELEGATIONS = [
|
||||
("moe_tp_rank", f"{_PS}.get_moe_tensor_parallel_rank"),
|
||||
("attn_tp_rank", f"{_PS}.get_attn_tensor_model_parallel_rank"),
|
||||
("attn_cp_rank", f"{_PS}.get_attn_context_model_parallel_rank"),
|
||||
("attn_dp_rank", f"{_DP}.get_attention_dp_rank"),
|
||||
]
|
||||
|
||||
GROUP_DELEGATIONS = [
|
||||
@@ -150,6 +151,87 @@ class TestParallelDelegation(_IsolatedOverrides):
|
||||
self.assertFalse(hasattr(ParallelContext, "local_attn_dp_size"))
|
||||
|
||||
|
||||
class TestStampedRanks(_IsolatedOverrides):
|
||||
"""`attn_dp_rank` comes from the stamp, and says so when there is none.
|
||||
|
||||
It is the one rank no group answers with: `initialize_dp_attention`
|
||||
computes it from this process's `tp_rank`, and an elastic scale-up
|
||||
replaces it with a rank in the expanded WORLD. Falling back to anything
|
||||
would be inventing a placement for this process.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
parallel = get_parallel()
|
||||
self._saved_derived = dict(parallel._derived)
|
||||
parallel.clear_derived_widths()
|
||||
self.addCleanup(
|
||||
lambda: (
|
||||
parallel.clear_derived_widths(),
|
||||
parallel.override_permanently(**self._saved_derived),
|
||||
)
|
||||
)
|
||||
|
||||
def test_the_stamp_is_the_answer(self):
|
||||
parallel = get_parallel()
|
||||
parallel.override_permanently(attn_dp_rank=3)
|
||||
self.assertEqual(parallel.attn_dp_rank, 3)
|
||||
# An elastic scale-up restamps it; the newest stamp wins.
|
||||
parallel.override_permanently(attn_dp_rank=9)
|
||||
self.assertEqual(parallel.attn_dp_rank, 9)
|
||||
|
||||
def test_a_scope_still_wins_over_the_stamp(self):
|
||||
parallel = get_parallel()
|
||||
parallel.override_permanently(attn_dp_rank=3)
|
||||
with parallel.override(attn_dp_rank=0):
|
||||
self.assertEqual(parallel.attn_dp_rank, 0)
|
||||
self.assertEqual(parallel.attn_dp_rank, 3)
|
||||
|
||||
def test_unstamped_names_the_cause(self):
|
||||
with self.assertRaises(RuntimeError) as caught:
|
||||
get_parallel().attn_dp_rank
|
||||
self.assertIn("initialize_dp_attention", str(caught.exception))
|
||||
|
||||
def test_a_stated_width_reaches_the_padding_mode(self):
|
||||
"""The reason this PR exists, from a reader's side.
|
||||
|
||||
`get_dp_padding_mode` reads the attention-DP width. Before the width
|
||||
had one home, a scoped `override` moved the context and left the
|
||||
module global answering, so stating a topology moved only half the
|
||||
runtime: this asserted `SUM_LEN` with the width stated as 1.
|
||||
"""
|
||||
from sglang.srt.layers.dp_attention import DpPaddingMode
|
||||
|
||||
with get_parallel().override(attn_dp_size=1):
|
||||
mode = DpPaddingMode.get_dp_padding_mode(
|
||||
is_extend_in_batch=True, global_num_tokens=[3, 5]
|
||||
)
|
||||
self.assertIs(mode, DpPaddingMode.MAX_LEN)
|
||||
|
||||
# And the branch it would have taken with the target's width.
|
||||
with get_parallel().override(attn_dp_size=2):
|
||||
mode = DpPaddingMode.get_dp_padding_mode(
|
||||
is_extend_in_batch=True, global_num_tokens=[3, 5]
|
||||
)
|
||||
self.assertIs(mode, DpPaddingMode.SUM_LEN)
|
||||
|
||||
def test_a_scale_up_stamps_the_width_and_the_rank_together(self):
|
||||
"""The two describe one topology; a reader that saw only one moved
|
||||
would place this process in a group it is not in."""
|
||||
from sglang.srt.layers.dp_attention import update_dp_attention_post_scale
|
||||
|
||||
# It also flips a process-wide gather flag; put it back, or every
|
||||
# later test in this process runs as if a scale-up had happened.
|
||||
dp_flags = get_flags().dp
|
||||
saved_gather = dp_flags.use_world_group_for_gather
|
||||
self.addCleanup(setattr, dp_flags, "use_world_group_for_gather", saved_gather)
|
||||
|
||||
parallel = get_parallel()
|
||||
update_dp_attention_post_scale(new_dp_size=16, new_dp_rank=11)
|
||||
self.assertEqual(parallel.attn_dp_size, 16)
|
||||
self.assertEqual(parallel.attn_dp_rank, 11)
|
||||
|
||||
|
||||
class TestParallelOverride(_IsolatedOverrides):
|
||||
def test_override_takes_precedence(self):
|
||||
p = get_parallel()
|
||||
@@ -1573,22 +1655,6 @@ class TestDerivedWidths(_IsolatedOverrides):
|
||||
with self.assertRaisesRegex(RuntimeError, r"derived parallel width"):
|
||||
get_parallel().attn_tp_size
|
||||
|
||||
def test_a_temporary_disable_beats_the_permanent_override(self):
|
||||
"""`disable_dp_size()` runs a draft scope without DP attention. It moves
|
||||
the module global the legacy getter reads, so it has to move the derived
|
||||
width too -- the scoped override wins over the permanent one, and a
|
||||
scope that left it alone would answer with the target model's width
|
||||
for its duration."""
|
||||
from sglang.srt.layers import dp_attention
|
||||
|
||||
parallel = get_parallel()
|
||||
parallel.override_permanently(attn_dp_size=4)
|
||||
with patch.object(dp_attention, "_ATTN_DP_SIZE", 4):
|
||||
with dp_attention.disable_dp_size():
|
||||
self.assertEqual(dp_attention.get_attention_dp_size(), 1)
|
||||
self.assertEqual(parallel.attn_dp_size, 1)
|
||||
self.assertEqual(parallel.attn_dp_size, 4)
|
||||
|
||||
def test_the_permanent_override_is_cleared_and_reset(self):
|
||||
parallel = get_parallel()
|
||||
parallel.override_permanently(attn_dp_size=2)
|
||||
|
||||
Reference in New Issue
Block a user