[Spec][PP] Launch extend microbatches before the spec output exchange (#40499)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user