[Spec][PP] Launch extend microbatches before the spec output exchange (#40499)

This commit is contained in:
YAMY
2026-09-21 15:47:03 -07:00
committed by GitHub
parent 0c53fec476
commit 0229025127
3 changed files with 47 additions and 13 deletions
@@ -17,7 +17,11 @@ maybe_stub_sgl_kernel()
from sglang.srt.managers.scheduler_components.request_receiver import ( # noqa: E402
SchedulerRequestReceiver,
)
from sglang.srt.managers.scheduler_pp_mixin import SchedulerPPMixin # noqa: E402
from sglang.srt.managers.scheduler_pp_mixin import ( # noqa: E402
SchedulerPPMixin,
_pp_exchange_outputs_before_forward,
)
from sglang.srt.model_executor.forward_batch_info import ForwardMode # noqa: E402
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
@@ -265,5 +269,20 @@ class TestDSparkPPOutput(CustomTestCase):
self.assertEqual(payloads[0].hidden_states.numel(), 0)
class TestPPSpecExchangeOrder(unittest.TestCase):
def test_extend_launches_before_the_relay_exchange(self):
kwargs = dict(spec_relay=True, is_last_rank=False, async_batch_depth=0)
extend = SimpleNamespace(
forward_mode=ForwardMode.EXTEND, is_extend_in_batch=False
)
decode = SimpleNamespace(
forward_mode=ForwardMode.DECODE, is_extend_in_batch=False
)
self.assertFalse(
_pp_exchange_outputs_before_forward(cur_batch=extend, **kwargs)
)
self.assertTrue(_pp_exchange_outputs_before_forward(cur_batch=decode, **kwargs))
if __name__ == "__main__":
unittest.main()