diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 9d5a96b85..c31c94a39 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -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 diff --git a/test/registered/distributed/test_dp_attention.py b/test/registered/distributed/test_dp_attention.py index c93ec10ce..53d86336a 100644 --- a/test/registered/distributed/test_dp_attention.py +++ b/test/registered/distributed/test_dp_attention.py @@ -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,