From 5ddca6819ebab6896b315967915161c1cb6b85e6 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:56:18 -0700 Subject: [PATCH] Fix unified SWA: size a non-owner's v2p by the id space it must address (#37560) Co-authored-by: Claude Opus 5 --- .../srt/mem_cache/multi_ended_allocator.py | 19 ++- .../test_gemma4_unified_swa_virtual_ids.py | 100 +++++++++++++++ .../test_unified_swa_shared_virtual_ids.py | 119 ++++++++++++++++++ 3 files changed, 232 insertions(+), 6 deletions(-) create mode 100644 test/registered/attention/test_gemma4_unified_swa_virtual_ids.py create mode 100644 test/registered/unit/mem_cache/test_unified_swa_shared_virtual_ids.py diff --git a/python/sglang/srt/mem_cache/multi_ended_allocator.py b/python/sglang/srt/mem_cache/multi_ended_allocator.py index 3d8ce76c6..6d0dfdc54 100644 --- a/python/sglang/srt/mem_cache/multi_ended_allocator.py +++ b/python/sglang/srt/mem_cache/multi_ended_allocator.py @@ -267,6 +267,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): sub_pool_name: str, device: str, is_id_owner: bool, + virtual_num_pages: Optional[int] = None, page_size: int = 1, shards_under_dcp: bool = False, need_sort: bool = False, @@ -332,10 +333,15 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): ) // self.pool_page_size self.entry_bytes_per_page = self.entry_bytes * self.pool_page_size - # v2p / p2v sized by PAGES. Page 0 is the padding anchor; trailing row is - # the -1 sentinel. + # v2p is indexed by VIRTUAL page id, p2v by PHYSICAL page id. A + # non-owner consumes the owner's ids, so its v2p spans the owner's + # count; the two are unrelated and either can be the larger. + self.num_virtual_ids = ( + self.num_pages if virtual_num_pages is None else virtual_num_pages + ) + # Page 0 is the padding anchor; the trailing row is the -1 sentinel. self.virtual_to_physical = torch.full( - (self.num_pages + 1,), + (self.num_virtual_ids + 1,), -1, dtype=torch.int64, device=device, @@ -346,8 +352,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): dtype=torch.int64, device=device, ) - # Back-compat alias (count of virtual PAGES) consulted by is_slot_allocated. - self.num_virtual_ids = self.num_pages # Chain neighbours: `low_peer` toward byte 0, `high_peer` toward # `total_bytes`. Ends have one (`bind_peer`), float middles have both. @@ -540,7 +544,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator): def is_slot_allocated(self, slot: int) -> bool: """Whether the PAGE containing this virtual id is in use.""" virt_page = slot // self.page_size - if virt_page < 0 or virt_page >= self.num_pages: + if virt_page < 0 or virt_page >= self.num_virtual_ids: return False return int(self.virtual_to_physical[virt_page].item()) != -1 @@ -3212,6 +3216,9 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator): need_sort=need_sort, forward_stream=forward_stream, lazy_compaction=lazy_compaction, + # swa binds the virtual pages full mints, so it must address + # full's whole id space. + virtual_num_pages=self.full_attn_allocator.num_virtual_ids, ) self._wire_peers() diff --git a/test/registered/attention/test_gemma4_unified_swa_virtual_ids.py b/test/registered/attention/test_gemma4_unified_swa_virtual_ids.py new file mode 100644 index 000000000..aee45cef7 --- /dev/null +++ b/test/registered/attention/test_gemma4_unified_swa_virtual_ids.py @@ -0,0 +1,100 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Unified memory on a hybrid sliding-window model, which nothing else covered. + +The unified SWA composite mints one virtual page id per allocation and binds it +on both sides, so the swa sub-pool's `virtual_to_physical` is indexed by the +FULL side's ids while it used to be sized by its own page count. A model with +few full-attention layers and many sliding ones gives the id owner more pages, +and the top of that id space then lands off the end of the swa table. On GPU +the write is unchecked (`alloc_bind_inplace`'s `tl.store`), so the symptom +surfaces later and elsewhere: a device-side index assert on the read in +`_swa_write_loc_unified`, and a dead scheduler. + +The reachability of it is a KV-budget property, not just a model property: ids +come off the front of the owner's free list and freed ids return to the back, so +the cursor sweeps the owner's range over time and the failure needs cumulative +churn past `swa.num_pages`. `--max-total-tokens 60000` puts that within one +200-example GSM8K run; at this model's default budget the swa table is 3.8M +entries and it would take dozens. + +Every argument below is load-bearing: a larger pool, or synthetic prompts in +place of this eval, and the narrow table passes. + + python -m pytest test/registered/attention/test_gemma4_unified_swa_virtual_ids.py -v +""" + +import types +import unittest + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.run_eval import run_eval +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=65, stage="base-b", runner_config="1-gpu-large") + +# Healthy 0.865-0.880 (static pool 0.870); a died-mid-run server scores +# 0.05-0.21. 1 sigma over 200 examples is ~0.024. +SCORE_THRESHOLD = 0.82 + + +class TestGemma4UnifiedSwaVirtualIds(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = "google/gemma-4-E2B-it" + 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=[ + "--attention-backend", + "triton", + "--enable-unified-memory", + "--disable-radix-cache", + "--mem-fraction-static", + "0.8", + # The trigger: this is what brings `swa.num_pages` within one + # eval's worth of churn. + "--max-total-tokens", + "60000", + ], + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + def test_gsm8k_survives_churn_past_the_swa_page_count(self): + metrics = run_eval( + types.SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + num_examples=200, + num_threads=128, + ) + ) + self.assertGreaterEqual(metrics["score"], SCORE_THRESHOLD) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_swa_shared_virtual_ids.py b/test/registered/unit/mem_cache/test_unified_swa_shared_virtual_ids.py new file mode 100644 index 000000000..8a4d45e7a --- /dev/null +++ b/test/registered/unit/mem_cache/test_unified_swa_shared_virtual_ids.py @@ -0,0 +1,119 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""The unified SWA composite mints one virtual page id and binds it on both +sides, so the swa side's `virtual_to_physical` is indexed by the FULL side's +ids. Its own page count says nothing about how wide that has to be: which side +gets more pages out of a shared byte budget depends on the layer split, and a +model with few full-attention layers and many sliding ones (gemma-4: 10 and 50) +gives the full side an order of magnitude more. + +Reaching an id that high takes cumulative allocation, so a narrow table fails +after churn rather than at once -- and on GPU it fails as a `tl.store` past the +end of the table: an unchecked write, not a raised index. + + python -m pytest test/registered/unit/mem_cache/test_unified_swa_shared_virtual_ids.py -v +""" + +import unittest + +import torch +from test_swa_locked_full_recover_unified import _DEV, _FakeUnifiedSWAKVPool + +from sglang.srt.mem_cache.multi_ended_allocator import ( + UnifiedSWATokenToKVPoolAllocator, +) +from sglang.srt.mem_cache.unified_memory_pool import MHASubPoolSpec, UnifiedKVPool +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +def _build(n_full: int, n_swa: int, full_layers: int, swa_layers: int): + """A composite whose sides have different per-page byte costs. `full_layers + < swa_layers` is the gemma-4 shape: the cheap side is the id owner and ends + up with far more pages than the side that has to address them.""" + full_spec = MHASubPoolSpec( + name="full", + layer_num=full_layers, + head_num=2, + head_dim=4, + store_dtype=torch.float16, + grow_direction="up", + ) + swa_spec = MHASubPoolSpec( + name="swa", + layer_num=swa_layers, + head_num=2, + head_dim=4, + store_dtype=torch.float16, + grow_direction="down", + ) + pool = UnifiedKVPool( + total_bytes=n_full * full_spec.entry_bytes() + n_swa * swa_spec.entry_bytes(), + sub_pool_specs=[full_spec, swa_spec], + device=_DEV, + enable_memory_saver=False, + ) + return UnifiedSWATokenToKVPoolAllocator( + unified_buffer=pool, + kvcache=_FakeUnifiedSWAKVPool(pool), + device=_DEV, + full_max_total_num_tokens=n_full, + swa_max_total_num_tokens=n_swa, + need_sort=False, + forward_stream=None, + ) + + +class TestSharedVirtualIdSpace(unittest.TestCase): + def test_swa_table_spans_the_owners_id_space(self): + """Static form: the table has to be wide enough before any alloc runs.""" + for full_layers, swa_layers in ((1, 5), (5, 1), (2, 2)): + with self.subTest(full_layers=full_layers, swa_layers=swa_layers): + alloc = _build(200, 20, full_layers, swa_layers) + owner = alloc.full_attn_allocator + swa = alloc.swa_attn_allocator + self.assertGreaterEqual( + int(swa.virtual_to_physical.shape[0]), + owner.num_virtual_ids + 1, + "swa v2p cannot address every id the owner can mint", + ) + # p2v stays this pool's own business: it is indexed by physical id. + self.assertEqual( + int(swa.physical_to_virtual.shape[0]), swa.num_pages + 1 + ) + + def test_churn_past_the_swa_page_count_binds_cleanly(self): + """Dynamic form: alloc/free until the owner's cursor passes the swa + side's page count, which is where the narrow table used to be written + off the end.""" + alloc = _build(200, 20, full_layers=1, swa_layers=5) + swa_pages = alloc.swa_attn_allocator.num_pages + highest = 0 + for _ in range(40): + v = alloc.alloc(4) + if v is None: + break + highest = max(highest, int(v.max()) // alloc.page_size) + alloc.free(v) + self.assertGreater( + highest, + swa_pages, + f"churn never reached past the swa side's {swa_pages} pages, so this " + "test would pass on a narrow table too", + ) + + +if __name__ == "__main__": + unittest.main()