[bugfix] Support MIXED forward mode in TBO splitter for DP attention (#24241)

This commit is contained in:
Cheng Wan
2026-05-01 16:01:23 -07:00
committed by GitHub
parent 05de73efd1
commit b47fab6f5d
2 changed files with 34 additions and 2 deletions
@@ -81,7 +81,7 @@ def compute_split_seq_index(
extend_lens: Optional[Sequence[int]],
token_num_per_seq: Optional[int],
) -> Optional[int]:
if forward_mode == ForwardMode.EXTEND:
if forward_mode == ForwardMode.EXTEND or forward_mode == ForwardMode.MIXED:
assert extend_lens is not None
return _split_extend_seqs(extend_lens)
elif forward_mode.is_target_verify() or forward_mode.is_decode():
@@ -270,7 +270,7 @@ def compute_split_token_index(
extend_seq_lens: Optional[Sequence[int]],
token_num_per_seq: Optional[int],
) -> int:
if forward_mode == ForwardMode.EXTEND:
if forward_mode == ForwardMode.EXTEND or forward_mode == ForwardMode.MIXED:
assert extend_seq_lens is not None
if _is_two_chunk_split_enabled(extend_seq_lens):
return sum(extend_seq_lens) // 2
@@ -68,6 +68,38 @@ class TestDPAttentionDP2TP2(
cls._env_override.__exit__(None, None, None)
class TestDPAttentionMixedChunk(
CustomTestCase,
GSM8KMixin,
):
gsm8k_accuracy_thres = 0.6
@classmethod
def setUpClass(cls):
cls.model = DEFAULT_MLA_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--trust-remote-code",
"--tp",
"2",
"--enable-dp-attention",
"--dp",
"2",
"--enable-mixed-chunk",
"--chunked-prefill-size",
"256",
],
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
class TestDPRetract(
CustomTestCase,
JSONConstrainedMixin,