[bugfix] Support MIXED forward mode in TBO splitter for DP attention (#24241)
This commit is contained in:
@@ -81,7 +81,7 @@ def compute_split_seq_index(
|
|||||||
extend_lens: Optional[Sequence[int]],
|
extend_lens: Optional[Sequence[int]],
|
||||||
token_num_per_seq: Optional[int],
|
token_num_per_seq: Optional[int],
|
||||||
) -> 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
|
assert extend_lens is not None
|
||||||
return _split_extend_seqs(extend_lens)
|
return _split_extend_seqs(extend_lens)
|
||||||
elif forward_mode.is_target_verify() or forward_mode.is_decode():
|
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]],
|
extend_seq_lens: Optional[Sequence[int]],
|
||||||
token_num_per_seq: Optional[int],
|
token_num_per_seq: Optional[int],
|
||||||
) -> int:
|
) -> int:
|
||||||
if forward_mode == ForwardMode.EXTEND:
|
if forward_mode == ForwardMode.EXTEND or forward_mode == ForwardMode.MIXED:
|
||||||
assert extend_seq_lens is not None
|
assert extend_seq_lens is not None
|
||||||
if _is_two_chunk_split_enabled(extend_seq_lens):
|
if _is_two_chunk_split_enabled(extend_seq_lens):
|
||||||
return sum(extend_seq_lens) // 2
|
return sum(extend_seq_lens) // 2
|
||||||
|
|||||||
@@ -68,6 +68,38 @@ class TestDPAttentionDP2TP2(
|
|||||||
cls._env_override.__exit__(None, None, None)
|
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(
|
class TestDPRetract(
|
||||||
CustomTestCase,
|
CustomTestCase,
|
||||||
JSONConstrainedMixin,
|
JSONConstrainedMixin,
|
||||||
|
|||||||
Reference in New Issue
Block a user