Keep CP current IPC staging proportional to touched pages
Cache-hit bs>1 current reuse can create very large dense attention buffers while touching only a small set of current pages. The previous SGLang runtime asked tai-kernel for a staging buffer sized like the full dense tensor, which caused CUDA OOM before the current IPC fast path could run. Switch token and index current IPC helpers to descriptor-compact staging: publish the dense destination pages into compact staging slots and materialize peers from compact source page ids back to the original dense destination pages. Document the failure mode and the compact-staging contract so the dense-sized contract is not reintroduced. Constraint: CUDA + SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE=1 must fail fast instead of silently falling back to current-slot all_reduce Rejected: Let staging allocation failure fall back to all_reduce | hides the bug and restores the expensive collective path Rejected: Size staging by the full dense tensor | reproduces the 965MB staging OOM on long-prefix cache-hit batches Confidence: high Scope-risk: moderate Directive: Current IPC helper source ids are compact staging ids; destination ids remain dense slot pages Tested: Remote cjy-glm5-new PYTHONPATH=python:/mnt/beegfs/cjy/tai-kernel/python python -m pytest -q test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -> 144 passed, 2 subtests passed Tested: Local py_compile cp_shared_kv_runtime.py Not-tested: Full ETE service restart with production traffic after this commit (cherry picked from commit 906ecbe5d4f08b73242e98e2b628e26516d5b04a)
This commit is contained in:
@@ -2874,7 +2874,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
self.assertEqual(mixed_locs.tolist(), [[4, 5]])
|
||||
self.assertTrue(torch.equal(mixed_kv[4:6], current_kv))
|
||||
|
||||
def test_current_token_ipc_helper_uses_dense_slot_pages_for_staging(self):
|
||||
def test_current_token_ipc_helper_uses_compact_pages_for_staging(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
@@ -2891,7 +2891,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
calls = []
|
||||
|
||||
class FakeKernels:
|
||||
def publish_cuda_ipc_slot_pages_and_mark_ready(self, *args, **kwargs):
|
||||
def publish_cuda_ipc_slot_pages_compact_and_mark_ready(self, *args, **kwargs):
|
||||
calls.append(("publish", args, kwargs))
|
||||
|
||||
def materialize_cuda_ipc_peer_pages_slot_indices_wait_ready(
|
||||
@@ -2922,7 +2922,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
gather_name, gather_args, gather_kwargs = calls[1]
|
||||
self.assertEqual(gather_name, "gather")
|
||||
self.assertTrue(torch.equal(gather_args[3], torch.tensor([0, 1])))
|
||||
self.assertTrue(torch.equal(gather_args[4], torch.tensor([1, 2])))
|
||||
self.assertTrue(torch.equal(gather_args[4], torch.tensor([0, 1])))
|
||||
self.assertTrue(torch.equal(gather_args[5], torch.tensor([1, 2])))
|
||||
self.assertEqual(gather_kwargs["ready_seq"], 1)
|
||||
|
||||
@@ -3463,7 +3463,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
self.assertTrue(torch.equal(dense_page_buffer[1, 0:4], current_k[0]))
|
||||
self.assertTrue(torch.equal(dense_page_buffer[1, 4:8], current_k[1]))
|
||||
|
||||
def test_current_index_ipc_helper_uses_dense_slot_pages_for_staging(self):
|
||||
def test_current_index_ipc_helper_uses_compact_pages_for_staging(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
@@ -3480,7 +3480,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
calls = []
|
||||
|
||||
class FakeKernels:
|
||||
def publish_cuda_ipc_slot_pages_and_mark_ready(self, *args, **kwargs):
|
||||
def publish_cuda_ipc_slot_pages_compact_and_mark_ready(self, *args, **kwargs):
|
||||
calls.append(("publish", args, kwargs))
|
||||
|
||||
def materialize_cuda_ipc_peer_pages_slot_indices_wait_ready(
|
||||
@@ -3505,7 +3505,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
self.assertTrue(torch.equal(calls[0][1][2], torch.tensor([2, 3])))
|
||||
self.assertEqual(calls[0][2]["ready_seq"], 7)
|
||||
self.assertTrue(torch.equal(calls[1][1][3], torch.tensor([1, -1])))
|
||||
self.assertTrue(torch.equal(calls[1][1][4], torch.tensor([2, -1])))
|
||||
self.assertTrue(torch.equal(calls[1][1][4], torch.tensor([0, -1])))
|
||||
self.assertTrue(torch.equal(calls[1][1][5], torch.tensor([2, 3])))
|
||||
self.assertEqual(calls[1][2]["ready_seq"], 7)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user