Replace skip_attn_backend_init with a batch-carried attention plan marker (+ staleness re-plan) (#27193)
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
47377525cb
commit
0aa72a9e76
@@ -0,0 +1,71 @@
|
||||
"""Regression: TBO filter_batch resets the attention plan marker on children.
|
||||
|
||||
filter_batch's completeness guard raises for any non-None ForwardBatch field
|
||||
missing from the child dict; the plan marker defaults to False (non-None) and
|
||||
crashed TBO cuda-graph capture until reset. CPU-only.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
import sglang.srt.batch_overlap.two_batch_overlap as tbo
|
||||
from sglang.srt.batch_overlap.two_batch_overlap import TboForwardBatchPreparer
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _make_target_verify_batch(bs: int) -> ForwardBatch:
|
||||
return ForwardBatch(
|
||||
forward_mode=ForwardMode.TARGET_VERIFY,
|
||||
batch_size=bs,
|
||||
input_ids=torch.zeros(bs, dtype=torch.long),
|
||||
positions=torch.zeros(bs, dtype=torch.long),
|
||||
out_cache_loc=torch.zeros(bs, dtype=torch.long),
|
||||
req_pool_indices=torch.zeros(bs, dtype=torch.long),
|
||||
seq_lens=torch.ones(bs, dtype=torch.int32),
|
||||
seq_lens_cpu=torch.ones(bs, dtype=torch.int32),
|
||||
seq_lens_sum=bs,
|
||||
spec_info=None,
|
||||
)
|
||||
|
||||
|
||||
def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch:
|
||||
fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3")
|
||||
with patch.object(tbo, "get_attention_tp_size", lambda: 1), patch.object(
|
||||
tbo, "get_global_server_args", lambda: fake_args
|
||||
):
|
||||
return TboForwardBatchPreparer.filter_batch(
|
||||
batch,
|
||||
start_token_index=lo,
|
||||
end_token_index=hi,
|
||||
start_seq_index=lo,
|
||||
end_seq_index=hi,
|
||||
out_num_token_non_padded=torch.tensor(hi - lo),
|
||||
)
|
||||
|
||||
|
||||
class TestTboFilterBatchMarker(CustomTestCase):
|
||||
def test_filter_batch_resets_plan_marker_on_children(self):
|
||||
child = _filter(_make_target_verify_batch(8), lo=0, hi=4)
|
||||
self.assertEqual(child.batch_size, 4)
|
||||
self.assertFalse(child.forward_metadata_ready)
|
||||
self.assertIsNone(child.forward_metadata_planned_bs)
|
||||
self.assertIsNone(child.forward_metadata_planned_num_tokens)
|
||||
self.assertFalse(child.forward_metadata_replan_equivalent)
|
||||
|
||||
def test_pre_planned_parent_does_not_leak_ready_into_children(self):
|
||||
parent = _make_target_verify_batch(8)
|
||||
parent.mark_forward_metadata_ready(replan_equivalent=True)
|
||||
child = _filter(parent, lo=0, hi=4)
|
||||
self.assertFalse(child.forward_metadata_ready)
|
||||
self.assertFalse(child.forward_metadata_replan_equivalent)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,127 @@
|
||||
"""Unit tests for the ForwardBatch attention plan marker / plan record.
|
||||
|
||||
Covers the contract behind ``skip_attn_backend_init`` deprecation:
|
||||
* fresh batches need planning; marked batches don't
|
||||
* the plan record (planned bs / num tokens) snapshots mark-time shapes
|
||||
* reshape after marking triggers a re-plan only for sites that opted
|
||||
into ``replan_equivalent``; re-marking re-records the new shapes
|
||||
* the deprecated kwarg shim maps explicit values onto the marker
|
||||
(mapped, not ignored) and warns once per process
|
||||
|
||||
Pure dataclass logic — CPU only.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
|
||||
import sglang.srt.model_executor.forward_batch_info as fbi
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _make_batch(bs: int = 2, num_tokens: int = 2) -> ForwardBatch:
|
||||
return ForwardBatch(
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
batch_size=bs,
|
||||
input_ids=torch.zeros(num_tokens, dtype=torch.long),
|
||||
req_pool_indices=torch.zeros(bs, dtype=torch.long),
|
||||
seq_lens=torch.ones(bs, dtype=torch.long),
|
||||
out_cache_loc=torch.zeros(bs, dtype=torch.long),
|
||||
seq_lens_sum=bs,
|
||||
)
|
||||
|
||||
|
||||
class TestForwardMetadataPlanRecord(CustomTestCase):
|
||||
def test_fresh_batch_needs_planning(self):
|
||||
fb = _make_batch()
|
||||
self.assertTrue(fb.needs_forward_metadata_init())
|
||||
self.assertFalse(fb.forward_metadata_ready)
|
||||
|
||||
def test_mark_records_shapes_and_skips_planning(self):
|
||||
fb = _make_batch(bs=3, num_tokens=7)
|
||||
fb.mark_forward_metadata_ready()
|
||||
self.assertFalse(fb.needs_forward_metadata_init())
|
||||
self.assertEqual(fb.forward_metadata_planned_bs, 3)
|
||||
self.assertEqual(fb.forward_metadata_planned_num_tokens, 7)
|
||||
|
||||
def test_reshape_without_opt_in_keeps_skipping(self):
|
||||
# Wrapper regimes must never auto-re-plan (would clobber per-step metadata).
|
||||
fb = _make_batch(bs=2)
|
||||
fb.mark_forward_metadata_ready()
|
||||
fb.batch_size = 4 # DP padding reshapes the batch
|
||||
self.assertFalse(fb.needs_forward_metadata_init())
|
||||
|
||||
def test_reshape_with_opt_in_replans(self):
|
||||
fb = _make_batch(bs=2, num_tokens=2)
|
||||
fb.mark_forward_metadata_ready(replan_equivalent=True)
|
||||
self.assertFalse(fb.needs_forward_metadata_init())
|
||||
|
||||
fb.batch_size = 4 # bs drift (prepare_mlp_sync_batch decode pad)
|
||||
self.assertTrue(fb.needs_forward_metadata_init())
|
||||
|
||||
fb.batch_size = 2
|
||||
fb.input_ids = torch.zeros(6, dtype=torch.long) # token drift
|
||||
self.assertTrue(fb.needs_forward_metadata_init())
|
||||
|
||||
def test_remark_re_records_padded_shapes(self):
|
||||
# Per-step loops re-mark each plan; the re-mark must snapshot padded shapes.
|
||||
fb = _make_batch(bs=2)
|
||||
fb.mark_forward_metadata_ready(replan_equivalent=True)
|
||||
fb.batch_size = 4
|
||||
self.assertTrue(fb.needs_forward_metadata_init())
|
||||
fb.mark_forward_metadata_ready(replan_equivalent=True)
|
||||
self.assertFalse(fb.needs_forward_metadata_init())
|
||||
self.assertEqual(fb.forward_metadata_planned_bs, 4)
|
||||
|
||||
|
||||
class TestDeprecatedSkipKwargShim(CustomTestCase):
|
||||
def setUp(self):
|
||||
self._saved_warned = fbi._skip_attn_backend_init_warned
|
||||
fbi._skip_attn_backend_init_warned = False
|
||||
|
||||
def tearDown(self):
|
||||
fbi._skip_attn_backend_init_warned = self._saved_warned
|
||||
|
||||
def test_none_is_a_silent_no_op(self):
|
||||
fb = _make_batch()
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("error")
|
||||
fb.apply_deprecated_skip_attn_backend_init(None)
|
||||
self.assertTrue(fb.needs_forward_metadata_init())
|
||||
|
||||
def test_true_maps_onto_marker_and_warns(self):
|
||||
# Mapped, not ignored: a no-op would silently re-plan multi-step metadata.
|
||||
fb = _make_batch()
|
||||
with warnings.catch_warnings(record=True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
fb.apply_deprecated_skip_attn_backend_init(True)
|
||||
self.assertFalse(fb.needs_forward_metadata_init())
|
||||
self.assertEqual(len(caught), 1)
|
||||
self.assertTrue(issubclass(caught[0].category, DeprecationWarning))
|
||||
|
||||
def test_false_warns_but_does_not_mark(self):
|
||||
fb = _make_batch()
|
||||
with warnings.catch_warnings(record=True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
fb.apply_deprecated_skip_attn_backend_init(False)
|
||||
self.assertTrue(fb.needs_forward_metadata_init())
|
||||
self.assertEqual(len(caught), 1)
|
||||
|
||||
def test_warns_once_per_process(self):
|
||||
# Hot-loop guard: per-forward callers must not pay warnings.warn repeatedly.
|
||||
fb = _make_batch()
|
||||
with warnings.catch_warnings(record=True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
fb.apply_deprecated_skip_attn_backend_init(True)
|
||||
fb.apply_deprecated_skip_attn_backend_init(True)
|
||||
_make_batch().apply_deprecated_skip_attn_backend_init(False)
|
||||
self.assertEqual(len(caught), 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user