[kv-shard 2/4] Sharded pools (#37615)

Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
Shunkangz
2026-09-16 19:22:23 +08:00
committed by GitHub
co-authored by Zhangheng
parent 76e06febab
commit e7f7447333
8 changed files with 1709 additions and 12 deletions
@@ -25,6 +25,7 @@ import torch
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.flashattention_backend import FlashAttentionBackend
from sglang.srt.mem_cache.kv_index_translator import KVIndexTranslator
from sglang.srt.mem_cache.page_interleave_pool import PageInterleaveKVPoolMixin
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
@@ -32,7 +33,13 @@ from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
def _make_prefill_aware_swa_runner(*, pool_size: int, max_context_len: int = 64):
class _FakeShardPool(PageInterleaveKVPoolMixin):
pass
def _make_prefill_aware_swa_runner(
*, pool_size: int, max_context_len: int = 64, token_to_kv_pool=None
):
"""A minimal fake ModelRunner that reaches FlashAttentionBackend.__init__'s
is_prefill_aware_swa branch (mirrors how models like
python/sglang/srt/models/unlimited_ocr.py opt in)."""
@@ -63,7 +70,7 @@ def _make_prefill_aware_swa_runner(*, pool_size: int, max_context_len: int = 64)
enable_prefill_cp=False,
enable_dp_attention=False,
)
token_to_kv_pool = object()
token_to_kv_pool = token_to_kv_pool if token_to_kv_pool is not None else object()
token_to_kv_pool_allocator = object()
return SimpleNamespace(
sliding_window_size=None,
@@ -94,6 +101,16 @@ def _make_prefill_aware_swa_runner(*, pool_size: int, max_context_len: int = 64)
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
class TestPrefillAwareSwaPrefillLensBound(CustomTestCase):
def test_sharded_pool_requests_cpu_sequence_lengths(self):
runner = _make_prefill_aware_swa_runner(
pool_size=8, token_to_kv_pool=_FakeShardPool()
)
with get_context().override_server_args():
backend = FlashAttentionBackend(runner)
self.assertTrue(backend.needs_cpu_seq_lens)
def test_buffer_covers_full_req_pool_idx_range(self):
pool_size = 8
runner = _make_prefill_aware_swa_runner(pool_size=pool_size)
@@ -0,0 +1,112 @@
# 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.
# ==============================================================================
import sys
from types import SimpleNamespace
import pytest
import torch
from sglang.srt.layers.attention.kv_shard_hooks import (
get_kv_shard_pool,
prepare_kv_shard_forward,
)
from sglang.srt.mem_cache.page_interleave_pool import PageInterleaveKVPoolMixin
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class _RecordingPool(PageInterleaveKVPoolMixin):
def __init__(self):
self.begin_args = None
self.begin_calls = 0
self.end_calls = 0
def begin_shard_extend(self, *args):
self.begin_args = args
self.begin_calls += 1
def end_shard_extend(self):
self.end_calls += 1
def _batch(mode: ForwardMode):
return SimpleNamespace(
forward_mode=mode,
req_pool_indices=torch.tensor([3], dtype=torch.int64),
extend_prefix_lens_cpu=[64],
seq_lens_cpu=torch.tensor([128], dtype=torch.int64),
)
def test_detects_only_page_interleaved_pools():
pool = _RecordingPool()
assert get_kv_shard_pool(pool) is pool
assert get_kv_shard_pool(object()) is None
@pytest.mark.parametrize(
"mode, active",
[
(ForwardMode.EXTEND, True),
(ForwardMode.MIXED, True),
(ForwardMode.SPLIT_PREFILL, True),
(ForwardMode.DECODE, False),
(ForwardMode.IDLE, False),
(ForwardMode.TARGET_VERIFY, False),
(ForwardMode.DRAFT_EXTEND_V2, False),
],
)
def test_prepare_updates_the_pool_lifecycle(mode, active):
pool = _RecordingPool()
req_to_token = torch.arange(8)
batch = _batch(mode)
assert prepare_kv_shard_forward(pool, req_to_token, batch) is active
assert pool.begin_calls == int(active)
assert pool.end_calls == int(not active)
if active:
assert all(
actual is expected
for actual, expected in zip(
pool.begin_args,
(
req_to_token,
batch.req_pool_indices,
batch.extend_prefix_lens_cpu,
batch.seq_lens_cpu,
),
)
)
else:
assert pool.begin_args is None
@pytest.mark.parametrize(
"missing_field",
["req_pool_indices", "extend_prefix_lens_cpu", "seq_lens_cpu"],
)
def test_extend_requires_host_metadata(missing_field):
batch = _batch(ForwardMode.EXTEND)
setattr(batch, missing_field, None)
with pytest.raises(RuntimeError, match="requires request indices and CPU"):
prepare_kv_shard_forward(_RecordingPool(), torch.empty(0), batch)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))
@@ -11,9 +11,10 @@
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Unit tests for logical-page KV cache sharding (CPU only).
"""Tests for logical-page KV cache sharding.
Pins the pure arithmetic that rotated owner-classed allocation hangs on:
Two sections. The first (CPU only, what the CPU CI job runs) pins the pure
arithmetic that rotated owner-classed allocation hangs on:
1. The placement bijection ``loc = Q*(N*ps) + r*ps + o`` — owner / local-row
round-trip, disjoint equal partition across ranks.
@@ -24,15 +25,33 @@ Pins the pure arithmetic that rotated owner-classed allocation hangs on:
3. The host rotation base on ``UnifiedTreeNode`` — stamped at insert, copied
on split, read through ``last_node``, and the pre-flight that declines an
insert whose pages carry a different base than the chain it would join.
4. ``translate_loc_to_scratch`` — the per-batch page->scratch-page lookup mapping
any consumer index vector onto the owner-major ``[prefix | chunk | trash]``
scratch, checked against a brute-force reference.
5. ``begin_shard_extend`` plan capture (page positions, padded send rows,
owner-congruence guard) with the gather stubbed out, following the
SimpleNamespace binding pattern of ``test_dsa_layer_shard_utils.py``.
The second section (``TestPageInterleaveGatherMultiGpu``, at the bottom) drives
real pools over a real 2-rank process group. It is the only check that the plan
the CPU stub validates actually addresses the bytes NCCL delivers, so it is
skipped rather than dropped when fewer than 2 CUDA devices are visible — which
is every run of the CPU suite this file is registered to.
"""
import os
import unittest
import unittest.mock
from array import array
from types import SimpleNamespace
import torch
import torch.multiprocessing as mp
from sglang.srt.distributed import (
init_distributed_environment,
initialize_model_parallel,
)
from sglang.srt.mem_cache.allocator.page_interleave import (
PageInterleavePoolAllocator,
page_interleave_shard_size,
@@ -40,19 +59,30 @@ from sglang.srt.mem_cache.allocator.page_interleave import (
from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictResult,
InsertParams,
MatchPrefixParams,
)
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import _evict_until_allocatable
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
from sglang.srt.mem_cache.page_interleave import (
PageInterleavePlacement,
PageShardSpec,
get_kv_shard_group,
)
from sglang.srt.mem_cache.page_interleave_pool import (
PageInterleaveKVPoolMixin,
PageInterleaveMHATokenToKVPool,
PageInterleaveMLATokenToKVPool,
)
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.unified_cache.components import ComponentType
from sglang.srt.mem_cache.unified_cache.unified_tree_core import UnifiedTreeCore
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
from sglang.srt.runtime_context import get_parallel, publish
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ceil_div
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -391,8 +421,6 @@ class TestEvictUntilAllocatable(CustomTestCase):
return alloc, outs
def _tree_stub(self, alloc, frees):
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
stub = SimpleNamespace(calls=0)
def evict(params):
@@ -407,8 +435,6 @@ class TestEvictUntilAllocatable(CustomTestCase):
return stub
def test_iterates_until_min_class_covers(self):
from sglang.srt.mem_cache.common import _evict_until_allocatable
alloc, outs = self._allocator_with_tight_class()
# Each round frees ONE class-3 page (a whole 1-page chain): reaching
# a min-class floor of 2 pages takes 2 rounds.
@@ -418,8 +444,6 @@ class TestEvictUntilAllocatable(CustomTestCase):
self.assertEqual(tree.calls, 2)
def test_terminates_when_tree_dry(self):
from sglang.srt.mem_cache.common import _evict_until_allocatable
alloc, _ = self._allocator_with_tight_class()
tree = self._tree_stub(alloc, []) # nothing evictable
_evict_until_allocatable(tree, alloc, PS)
@@ -767,5 +791,707 @@ class TestRotationGraftDecline(CustomTestCase):
self.assertEqual(set(released.tolist()), set(own_locs[:8].tolist()))
def _chain_pages(base, n_pages, local_start=5):
"""Logical page ids of one chain: page P has owner (base + P) % N and an
arbitrary (here: increasing) local page on its owner."""
counter = {r: local_start for r in range(N)}
pages = []
for p in range(n_pages):
r = (base + p) % N
pages.append(counter[r] * N + r)
counter[r] += 1
return pages
def _chain_row(pages, seq_len):
row = torch.empty(seq_len, dtype=torch.int32)
for i in range(seq_len):
row[i] = pages[i // PS] * PS + i % PS
return row
def _make_pool_stub(spec, shard_rank=0, debug=True, table_pages=4096):
"""A SimpleNamespace carrying exactly the state begin_shard_extend /
translate_loc_to_scratch read."""
stub = SimpleNamespace()
stub.shard_spec = spec
stub.shard_rank = shard_rank
stub.device = "cpu"
stub.start_layer = 0
stub._chunk_base = spec.max_prefix_tokens
stub._trash_base = spec.max_prefix_tokens + spec.chunk_tokens
stub._page_pos = torch.full(
(table_pages,), stub._trash_base // PS, dtype=torch.int32
)
stub._local_page_stride = table_pages
stub._epoch = 0
stub._write_plan_key = stub._write_plan = None
stub._translate_cache = {}
stub._debug_plan_checks = debug
stub.translate_loc_to_scratch = lambda loc: (
PageInterleaveKVPoolMixin.translate_loc_to_scratch(stub, loc)
)
stub.prefetched = []
stub._prefetch_layer = lambda layer_id: stub.prefetched.append(layer_id)
return stub
def _run_begin(stub, prefix_lens, seq_lens, rows):
width = max(r.numel() for r in rows)
padded = [
torch.cat([r, torch.zeros(width - r.numel(), dtype=torch.int32)]) for r in rows
]
PageInterleaveKVPoolMixin.begin_shard_extend(
stub,
torch.stack(padded),
torch.arange(len(rows)),
prefix_lens,
seq_lens,
)
return stub
def _reference_prefix_slots(per_request_prefix_pages):
"""Brute-force reference of the owner-major slot assignment: the batch's
unique prefix pages sorted by (owner, local page), rank r's pages
contiguous at r * block; block = sum of per-request ceil(K_i / N)."""
block = sum(ceil_div(len(pages), N) for pages in per_request_prefix_pages)
uniq = sorted({p for pages in per_request_prefix_pages for p in pages})
slots = {}
counts = {r: 0 for r in range(N)}
for page in sorted(uniq, key=lambda p: (p % N, p // N)):
owner = page % N
slots[page] = owner * block + counts[owner]
counts[owner] += 1
return slots, block
class TestBeginShardExtendPlan(CustomTestCase):
def test_plan_with_rotated_prefix(self):
"""7 prefix pages of a base-2 chain + 9 chunk pages (last partial):
owner-major slots, send rows owner-filtered in the same order and
padded to the block bound ceil(7/4) = 2 pages."""
pages = _chain_pages(base=2, n_pages=16)
prefix_len, seq_len = 7 * PS, 16 * PS - 5
row = _chain_row(pages, seq_len)
slots, block = _reference_prefix_slots([pages[:7]])
for rank in range(N):
stub = _run_begin(
_make_pool_stub(_make_spec(), rank), [prefix_len], [seq_len], [row]
)
self.assertEqual(stub._block_pages, block)
self.assertTrue(stub._shard_extend_active)
self.assertEqual(stub._epoch, 1)
self.assertEqual(stub.prefetched, [0]) # first layer kicked
for page, slot in slots.items():
self.assertEqual(int(stub._page_pos[page]), slot)
for j, page in enumerate(pages[7:]):
self.assertEqual(int(stub._page_pos[page]), stub._chunk_base // PS + j)
own = sorted((p for p in pages[:7] if p % N == rank), key=lambda p: p // N)
expect = torch.cat(
[torch.arange((p // N) * PS, (p // N + 1) * PS) for p in own]
)
if len(own) < block: # padded with the trash page (local page 0)
expect = torch.cat([expect, torch.arange((block - len(own)) * PS)])
self.assertTrue(torch.equal(stub._send_rows, expect))
def test_multi_request_plan_shared_prefix_dedup(self):
"""bs > 1: request 0 and request 1 share a 3-page cached prefix
(request 1 extends it by 2 pages); request 2 is an unrelated base-2
chain. Shared pages must gather into ONE slot (no duplicate plan
entries), the block is the per-request ceil sum, and every request's
locs translate through the same table."""
chain_a = _chain_pages(base=0, n_pages=5)
chain_c = _chain_pages(base=2, n_pages=4, local_start=20)
# rows: request 0 = A[:3] prefix + 1 chunk page; request 1 = A[:5]
# prefix + 2 chunk pages; request 2 = C[:2] prefix + 2 chunk pages.
chunk0 = _chain_pages(base=3, n_pages=1, local_start=40)
chunk1 = _chain_pages(base=1, n_pages=2, local_start=50)
chunk2 = _chain_pages(base=0, n_pages=2, local_start=60)
rows = [
_chain_row(chain_a[:3] + chunk0, 4 * PS),
_chain_row(chain_a[:5] + chunk1, 7 * PS),
_chain_row(chain_c[:2] + chunk2, 4 * PS - 3),
]
stub = _run_begin(
_make_pool_stub(_make_spec()),
[3 * PS, 5 * PS, 2 * PS],
[4 * PS, 7 * PS, 4 * PS - 3],
rows,
)
slots, block = _reference_prefix_slots([chain_a[:3], chain_a[:5], chain_c[:2]])
self.assertEqual(block, 1 + 2 + 1)
self.assertEqual(stub._block_pages, block)
for page, slot in slots.items():
self.assertEqual(int(stub._page_pos[page]), slot)
# Chunk slots are absolute scratch pages in batch order.
for j, page in enumerate(chunk0 + chunk1 + chunk2):
self.assertEqual(int(stub._page_pos[page]), stub._chunk_base // PS + j)
# Shared pages: both requests' locs hit the SAME scratch rows.
shared_loc_r0 = rows[0][:PS].long()
shared_loc_r1 = rows[1][:PS].long()
t0 = PageInterleaveKVPoolMixin.translate_loc_to_scratch(stub, shared_loc_r0)
t1 = PageInterleaveKVPoolMixin.translate_loc_to_scratch(stub, shared_loc_r1)
self.assertTrue(torch.equal(t0, t1))
# Per-rank send lists fit the block and pad with the trash page.
all_prefix = sorted(set(chain_a[:5] + chain_c[:2]))
for rank in range(N):
stub_r = _run_begin(
_make_pool_stub(_make_spec(), rank),
[3 * PS, 5 * PS, 2 * PS],
[4 * PS, 7 * PS, 4 * PS - 3],
rows,
)
own = sorted((p for p in all_prefix if p % N == rank), key=lambda p: p // N)
self.assertLessEqual(len(own), block)
self.assertEqual(stub_r._send_rows.numel(), block * PS)
expect_head = torch.cat(
[torch.arange((p // N) * PS, (p // N + 1) * PS) for p in own]
)
self.assertTrue(
torch.equal(stub_r._send_rows[: len(own) * PS], expect_head)
)
def test_send_order_follows_local_page_not_position(self):
"""A freed-and-reused page can give a chain a LOWER local page id at
a later position. Slot assignment and send packing must both order
by local page id (they only need to agree — a mismatch reads the
wrong rank rows)."""
# Owner-0 pages appear at positions 0 and 4 with locals 9 then 3.
pages = [9 * N + 0, 5 * N + 1, 5 * N + 2, 5 * N + 3, 3 * N + 0]
row = _chain_row(pages, 5 * PS)
stub = _run_begin(
_make_pool_stub(_make_spec(), shard_rank=0),
[5 * PS],
[5 * PS + PS],
[torch.cat([row, _chain_row([7 * N + 1], PS)])],
)
slots, block = _reference_prefix_slots([pages])
self.assertEqual(block, 2)
# local 3 gets owner-0's first slot although it sits at position 4.
self.assertEqual(int(stub._page_pos[3 * N + 0]), 0)
self.assertEqual(int(stub._page_pos[9 * N + 0]), 1)
expect = torch.cat(
[torch.arange(3 * PS, 4 * PS), torch.arange(9 * PS, 10 * PS)]
)
self.assertTrue(torch.equal(stub._send_rows, expect))
def test_plan_without_prefix(self):
pages = _chain_pages(base=0, n_pages=2)
stub = _run_begin(
_make_pool_stub(_make_spec()), [0], [PS + 5], [_chain_row(pages, PS + 5)]
)
self.assertEqual(stub._block_pages, 0)
self.assertTrue(stub._shard_extend_active)
self.assertEqual(stub.prefetched, []) # nothing to gather
self.assertIsNone(stub._send_rows)
self.assertEqual(int(stub._page_pos[pages[0]]), stub._chunk_base // PS)
self.assertEqual(int(stub._page_pos[pages[1]]), stub._chunk_base // PS + 1)
def test_unaligned_prefix_rejected(self):
# The tree quantum is the PHYSICAL page: a prefix that is not a
# ps-multiple can never come out of match_prefix.
pages = _chain_pages(base=0, n_pages=4)
with self.assertRaises(AssertionError):
_run_begin(
_make_pool_stub(_make_spec()),
[PS + 3],
[4 * PS],
[_chain_row(pages, 4 * PS)],
)
def test_owner_congruence_guard(self):
"""A rotation-base bug that breaks a request's prefix-owner
cyclicity invalidates the sync-free block bound (a rank can own more
than ceil(K/N) pages); the debug guard must catch it at plan time."""
pages = _chain_pages(base=1, n_pages=8)
pages[2], pages[5] = pages[5], pages[2] # same multiset, not cyclic
with self.assertRaises(AssertionError) as ctx:
_run_begin(
_make_pool_stub(_make_spec()),
[6 * PS],
[8 * PS],
[_chain_row(pages, 8 * PS)],
)
self.assertIn("cyclic", str(ctx.exception))
class TestScratchTranslation(CustomTestCase):
def _plan(self, base=2, n_prefix=7, n_chunk=9, rank=1):
pages = _chain_pages(base=base, n_pages=n_prefix + n_chunk)
seq_len = (n_prefix + n_chunk) * PS
stub = _run_begin(
_make_pool_stub(_make_spec(), rank),
[n_prefix * PS],
[seq_len],
[_chain_row(pages, seq_len)],
)
return stub, pages[:n_prefix], pages[n_prefix:]
def _reference_row(self, stub, prefix_pages, chunk_pages, loc):
"""Brute-force reference: owner-major (owner, local-page)-sorted
prefix slots, sequence-order chunk."""
spec = stub.shard_spec
page, off = loc // PS, loc % PS
if page in prefix_pages:
slots, _ = _reference_prefix_slots([prefix_pages])
return slots[page] * PS + off
if page in chunk_pages:
k = chunk_pages.index(page)
return spec.max_prefix_tokens + k * PS + off
return stub._trash_base + off
def test_translation_matches_reference(self):
stub, prefix_pages, chunk_pages = self._plan()
locs = (
[p * PS + o for p in prefix_pages + chunk_pages for o in (0, 3, PS - 1)]
+ list(range(0, N)) # reserved pages -> trash
+ [3000, 3001] # off-plan -> trash
)
got = PageInterleaveKVPoolMixin.translate_loc_to_scratch(
stub, torch.tensor(locs, dtype=torch.int64)
)
expect = torch.tensor(
[self._reference_row(stub, prefix_pages, chunk_pages, l) for l in locs],
dtype=torch.int64,
)
self.assertTrue(torch.equal(got, expect))
def test_translation_is_injective_over_the_plan(self):
stub, prefix_pages, chunk_pages = self._plan(base=3, n_prefix=5, n_chunk=4)
locs = [p * PS + o for p in prefix_pages + chunk_pages for o in range(PS)]
rows = PageInterleaveKVPoolMixin.translate_loc_to_scratch(
stub, torch.tensor(locs, dtype=torch.int64)
)
self.assertEqual(len(torch.unique(rows)), len(locs))
# Prefix rows stay inside the (padded) gather span, chunk rows inside
# the chunk region.
n_prefix_tokens = len(prefix_pages) * PS
self.assertTrue(
bool((rows[:n_prefix_tokens] < N * stub._block_pages * PS).all())
)
self.assertTrue(
bool(
(rows[n_prefix_tokens:] >= stub.shard_spec.max_prefix_tokens).all()
and (rows[n_prefix_tokens:] < stub._trash_base).all()
)
)
def test_int32_page_table_input(self):
stub, prefix_pages, chunk_pages = self._plan(base=0, n_prefix=4, n_chunk=1)
table = torch.tensor(
[prefix_pages[0] * PS, prefix_pages[1] * PS, chunk_pages[0] * PS, 0],
dtype=torch.int32,
)
rows = PageInterleaveKVPoolMixin.translate_loc_to_scratch(stub, table)
self.assertEqual(rows.dtype, torch.int64)
# Page-aligned inputs land on page-aligned scratch rows (the FA3
# stride-divide contract).
self.assertTrue(bool((rows[:3] % PS == 0).all()))
self.assertEqual(int(rows[3]), stub._trash_base)
def test_translation_cache_cleared_with_new_plan(self):
pages = _chain_pages(base=0, n_pages=2)
stub = _make_pool_stub(_make_spec())
# The first batch treats page 0 as part of the current chunk.
_run_begin(stub, [0], [PS], [_chain_row(pages[:1], PS)])
loc = _chain_row(pages[:1], PS).long()
first = PageInterleaveKVPoolMixin._translate_loc_cached(stub, loc)
again = PageInterleaveKVPoolMixin._translate_loc_cached(stub, loc)
self.assertIs(again, first)
# The next batch reuses the same loc tensor after page 0 becomes a
# cached prefix. Installing the new plan must discard the old mapping.
_run_begin(stub, [PS], [2 * PS], [_chain_row(pages, 2 * PS)])
fresh = PageInterleaveKVPoolMixin._translate_loc_cached(stub, loc)
self.assertIsNot(fresh, first)
self.assertFalse(torch.equal(fresh, first))
class TestWritePlan(CustomTestCase):
def test_owner_filter_cached_per_loc_tensor(self):
spec = _make_spec(shard_rank=2)
stub = SimpleNamespace()
stub.placement = PageInterleavePlacement(spec)
stub.shard_rank = 2
stub._epoch = 1
stub._write_plan_key = stub._write_plan = None
loc = torch.arange(5 * GS, 7 * GS) # two whole groups
owned_idx, local_rows = PageInterleaveKVPoolMixin._get_write_plan(stub, loc)
self.assertEqual(owned_idx.numel(), 2 * PS)
# Owned rows are ps-contiguous runs at [Q*ps, (Q+1)*ps).
self.assertTrue(
torch.equal(
local_rows,
torch.cat([torch.arange(5 * PS, 6 * PS), torch.arange(6 * PS, 7 * PS)]),
)
)
# Same tensor + same epoch -> cached (identity).
again = PageInterleaveKVPoolMixin._get_write_plan(stub, loc)
self.assertIs(again[0], owned_idx)
# Epoch bump invalidates.
stub._epoch = 2
fresh = PageInterleaveKVPoolMixin._get_write_plan(stub, loc)
self.assertIsNot(fresh[0], owned_idx)
def test_partial_tail_page_may_own_nothing(self):
spec = _make_spec(shard_rank=3)
stub = SimpleNamespace()
stub.placement = PageInterleavePlacement(spec)
stub.shard_rank = 3
stub._epoch = 1
stub._write_plan_key = stub._write_plan = None
# 10 tokens: all inside owner-0's page of the group.
loc = torch.arange(8 * GS, 8 * GS + 10)
owned_idx, local_rows = PageInterleaveKVPoolMixin._get_write_plan(stub, loc)
self.assertEqual(owned_idx.numel(), 0)
self.assertEqual(local_rows.numel(), 0)
# =============================================================================
# Multi-GPU: the real NCCL layer-ahead gather (2 GPUs).
#
# Everything above is pure arithmetic on a CPU stub. This section drives real
# pools over a real process group, which is the only check that the plan the
# stub validates actually addresses the bytes the collective delivers:
#
# 1. MLA pool sharded across the attention-TP group: replicated writes are
# owner-filtered into disjoint pool stripes; a later batch's chunked-prefix
# read (get_mla_kv_buffer) assembles the full prefix from all ranks via the
# layer-ahead NCCL allgather and must return the canonical bytes.
# 2. MHA pool sharded across the attention-CP group: the post-allgather full
# chunk is staged into the scratch chunk region and owner-persisted; a later
# batch reads prefix+chunk through the translated page table (the scratch),
# and the assembled rows must match the canonical bytes.
#
# Skipped unless 2 CUDA devices are visible, so it is inert on the CPU runner
# this file is registered to. Run it explicitly with:
# CUDA_VISIBLE_DEVICES=0,1 python3 test/registered/unit/mem_cache/\
# test_page_interleave_shard.py TestPageInterleaveGatherMultiGpu
# =============================================================================
_GATHER_WORLD = 2
_GATHER_LAYER_NUM = 4
_GATHER_PAGE_SIZE = 16
_GATHER_GRANULE = _GATHER_WORLD * _GATHER_PAGE_SIZE
_GATHER_SIZE = _GATHER_PAGE_SIZE * 64 # physical token slots per rank
_GATHER_KV_LORA_RANK = 128
_GATHER_QK_ROPE = 32
_GATHER_HEAD_NUM = 2
_GATHER_HEAD_DIM = 32
_GATHER_DTYPE = torch.bfloat16
def _mla_value(loc, dim):
"""Deterministic canonical latent value for logical slot ``loc``."""
loc = loc.to(torch.float32)
return (loc.unsqueeze(-1) + torch.arange(dim, device=loc.device) * 0.001).to(
_GATHER_DTYPE
)
def _dist_init(rank, world, port, attn_cp_size):
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = str(port)
os.environ["RANK"] = str(rank)
os.environ["WORLD_SIZE"] = str(world)
os.environ.setdefault("no_proxy", "127.0.0.1,localhost")
torch.cuda.set_device(rank)
init_distributed_environment(
world_size=world,
rank=rank,
local_rank=rank,
distributed_init_method=f"tcp://127.0.0.1:{port}",
backend="nccl",
)
# Publish the widths the groups below are about to be built at. The derived
# quotients (attn_tp_size, attn_dcp_size, ...) are projected from these
# leaves at publish; initialize_model_parallel no longer supplies them, and
# MLATokenToKVPool.set_mla_kv_buffer reads attn_dcp_size on the write path.
publish(
ServerArgs(model_path="dummy", tp_size=world, attn_cp_size=attn_cp_size),
role="scheduler",
)
initialize_model_parallel(
tensor_model_parallel_size=world,
attention_context_model_parallel_size=attn_cp_size,
)
def _gather_make_spec(shard_rank, max_prefix_groups=16, chunk_groups=4):
return PageShardSpec(
shard_rank=shard_rank,
shard_size=_GATHER_WORLD,
page_size=_GATHER_PAGE_SIZE,
max_prefix_tokens=max_prefix_groups * _GATHER_GRANULE,
chunk_tokens=chunk_groups * _GATHER_GRANULE,
)
def _fake_req_to_token(groups, seq_len, device):
"""req_to_token row where sequence group j is allocator group groups[j]."""
row = torch.zeros(
(1, len(groups) * _GATHER_GRANULE), dtype=torch.int32, device=device
)
for j, q in enumerate(groups):
row[0, j * _GATHER_GRANULE : (j + 1) * _GATHER_GRANULE] = torch.arange(
q * _GATHER_GRANULE,
(q + 1) * _GATHER_GRANULE,
dtype=torch.int32,
device=device,
)
return row[:, :seq_len] if seq_len < row.shape[1] else row
def _check(rank, name, got, expect, atol=0.0):
ok = torch.allclose(got.float(), expect.float(), atol=atol, rtol=0)
max_err = (got.float() - expect.float()).abs().max().item()
print(f"[rank {rank}] {name}: max_err={max_err:.6f} {'OK' if ok else 'FAIL'}")
assert ok, f"[rank {rank}] {name} mismatch (max_err={max_err})"
def _run_mla(rank, world, port):
_dist_init(rank, world, port, attn_cp_size=1)
group = get_parallel().attn_tp_group
assert group.world_size == world
# Topology-first shard-group selection: no CP here, so MLA falls back to
# the attn-TP axis, while GQA has no replicated axis (world_size 1).
assert get_kv_shard_group(use_mla_backend=True) is group
assert get_kv_shard_group(use_mla_backend=False).world_size == 1
spec = _gather_make_spec(shard_rank=group.rank_in_group)
pool = PageInterleaveMLATokenToKVPool(
_GATHER_SIZE,
page_size=_GATHER_PAGE_SIZE,
dtype=_GATHER_DTYPE,
kv_lora_rank=_GATHER_KV_LORA_RANK,
qk_rope_head_dim=_GATHER_QK_ROPE,
layer_num=_GATHER_LAYER_NUM,
device=f"cuda:{rank}",
enable_memory_saver=False,
start_layer=0,
end_layer=_GATHER_LAYER_NUM - 1,
shard_spec=spec,
shard_group=group,
)
device = pool.kv_buffer[0].device
# ---- chunk 1: replicated write, owner-filtered persist -----------------
# "Allocator" hands out fragmented groups (identical on every rank).
chunk1_groups = [5, 2, 9]
chunk1_locs = _fake_req_to_token(chunk1_groups, 3 * _GATHER_GRANULE, device)[
0
].long()
for layer_id in range(_GATHER_LAYER_NUM):
layer = SimpleNamespace(layer_id=layer_id)
vals = _mla_value(
chunk1_locs + layer_id * 1000, _GATHER_KV_LORA_RANK + _GATHER_QK_ROPE
)
pool.set_mla_kv_buffer(
layer,
chunk1_locs,
vals[:, :_GATHER_KV_LORA_RANK].unsqueeze(1),
vals[:, _GATHER_KV_LORA_RANK:].unsqueeze(1),
)
torch.cuda.synchronize()
torch.distributed.barrier()
# Pool holds only the owned stripe: group Q sits at local rows [Q*ps,(Q+1)*ps)
# on every rank, holding that rank's page of the group.
for q in chunk1_groups:
local_rows = torch.arange(
q * _GATHER_PAGE_SIZE, (q + 1) * _GATHER_PAGE_SIZE, device=device
)
owned_locs = (
q * _GATHER_GRANULE
+ group.rank_in_group * _GATHER_PAGE_SIZE
+ torch.arange(_GATHER_PAGE_SIZE, device=device)
)
got = pool.kv_buffer[0][local_rows, 0, :].view(_GATHER_DTYPE)
expect = _mla_value(owned_locs, _GATHER_KV_LORA_RANK + _GATHER_QK_ROPE)
_check(rank, f"mla owned stripe g{q}", got, expect)
# ---- chunk 2: prefix gather + staged chunk, both read styles -----------
seq_groups = chunk1_groups + [12] # one new chunk group
prefix_len = 3 * _GATHER_GRANULE
seq_len = prefix_len + _GATHER_GRANULE
req_to_token = _fake_req_to_token(seq_groups, seq_len, device)
chunk2_locs = req_to_token[0, prefix_len:seq_len].long()
pool.begin_shard_extend(req_to_token, torch.tensor([0]), [prefix_len], [seq_len])
for layer_id in range(_GATHER_LAYER_NUM):
layer = SimpleNamespace(layer_id=layer_id)
# Write the current chunk (stages it into the slot + persists the
# owned stripe), like the extend forward does before attention.
chunk_vals = _mla_value(
chunk2_locs + layer_id * 1000, _GATHER_KV_LORA_RANK + _GATHER_QK_ROPE
)
pool.set_mla_kv_buffer(
layer,
chunk2_locs,
chunk_vals[:, :_GATHER_KV_LORA_RANK].unsqueeze(1),
chunk_vals[:, _GATHER_KV_LORA_RANK:].unsqueeze(1),
)
# Chunked-prefix MHA style: fetch an arbitrary sub-range of the
# prefix through get_mla_kv_buffer.
sub = chunk1_locs[_GATHER_PAGE_SIZE // 2 : prefix_len - 3]
k_nope, k_rope = pool.get_mla_kv_buffer(layer, sub, _GATHER_DTYPE)
expect = _mla_value(
sub + layer_id * 1000, _GATHER_KV_LORA_RANK + _GATHER_QK_ROPE
)
_check(
rank,
f"mla prefix read l{layer_id}",
k_nope[:, 0, :],
expect[:, :_GATHER_KV_LORA_RANK],
)
_check(
rank,
f"mla prefix rope l{layer_id}",
k_rope[:, 0, :],
expect[:, _GATHER_KV_LORA_RANK:],
)
# Absorbed-MLA style (what MLA-under-CP uses): read [prefix | chunk]
# from get_key_buffer through the translated page table.
all_locs = req_to_token[0, :seq_len].long()
rows = pool.translate_loc_to_scratch(all_locs)
kv_scratch = pool.get_key_buffer(layer_id)
_check(
rank,
f"mla absorbed read l{layer_id}",
kv_scratch[rows, 0, :],
_mla_value(
all_locs + layer_id * 1000, _GATHER_KV_LORA_RANK + _GATHER_QK_ROPE
),
)
torch.distributed.barrier()
if rank == 0:
print("PASS: MLA page-interleave shard (attn-TP axis)")
def _run_mha(rank, world, port):
_dist_init(rank, world, port, attn_cp_size=world)
group = get_parallel().attn_cp_group
assert group.world_size == world
# Topology-first shard-group selection: with an active CP group, both
# GQA and MLA shard across CP (CP replicates KV for every attention
# type; the TP axis is only the no-CP MLA fallback).
assert get_kv_shard_group(use_mla_backend=False) is group
assert get_kv_shard_group(use_mla_backend=True) is group
spec = _gather_make_spec(shard_rank=group.rank_in_group)
pool = PageInterleaveMHATokenToKVPool(
_GATHER_SIZE,
page_size=_GATHER_PAGE_SIZE,
dtype=_GATHER_DTYPE,
head_num=_GATHER_HEAD_NUM,
head_dim=_GATHER_HEAD_DIM,
layer_num=_GATHER_LAYER_NUM,
device=f"cuda:{rank}",
enable_memory_saver=False,
start_layer=0,
end_layer=_GATHER_LAYER_NUM - 1,
enable_alt_stream=False,
shard_spec=spec,
shard_group=group,
)
device = pool.k_buffer[0].device
def kv_value(locs, layer_id, is_v):
base = locs.to(torch.float32) + layer_id * 1000 + (500000 if is_v else 0)
return (
base.view(-1, 1, 1)
+ torch.arange(_GATHER_HEAD_NUM, device=device).view(1, -1, 1) * 0.01
+ torch.arange(_GATHER_HEAD_DIM, device=device).view(1, 1, -1) * 0.0001
).to(_GATHER_DTYPE)
# ---- chunk 1 (prefix-less batch): stage + owner-persist ----------------
chunk1_groups = [7, 3]
chunk1_locs = _fake_req_to_token(chunk1_groups, 2 * _GATHER_GRANULE, device)[
0
].long()
req_to_token = _fake_req_to_token(chunk1_groups, 2 * _GATHER_GRANULE, device)
pool.begin_shard_extend(req_to_token, torch.tensor([0]), [0], [2 * _GATHER_GRANULE])
for layer_id in range(_GATHER_LAYER_NUM):
layer = SimpleNamespace(layer_id=layer_id)
pool.set_kv_buffer(
layer,
chunk1_locs,
kv_value(chunk1_locs, layer_id, False),
kv_value(chunk1_locs, layer_id, True),
)
# The current chunk must be readable through the scratch right away.
k_scratch = pool.get_key_buffer(layer_id)
rows = pool.translate_loc_to_scratch(chunk1_locs)
_check(
rank,
f"mha chunk stage l{layer_id}",
k_scratch[rows],
kv_value(chunk1_locs, layer_id, False),
)
torch.cuda.synchronize()
torch.distributed.barrier()
# ---- chunk 2: prefix gathered from peers via translated page table -----
seq_groups = chunk1_groups + [11]
prefix_len = 2 * _GATHER_GRANULE
seq_len = prefix_len + _GATHER_GRANULE
req_to_token = _fake_req_to_token(seq_groups, seq_len, device)
chunk2_locs = req_to_token[0, prefix_len:seq_len].long()
pool.begin_shard_extend(req_to_token, torch.tensor([0]), [prefix_len], [seq_len])
for layer_id in range(_GATHER_LAYER_NUM):
layer = SimpleNamespace(layer_id=layer_id)
pool.set_kv_buffer(
layer,
chunk2_locs,
kv_value(chunk2_locs, layer_id, False),
kv_value(chunk2_locs, layer_id, True),
)
all_locs = req_to_token[0, :seq_len].long()
rows = pool.translate_loc_to_scratch(all_locs)
k_scratch = pool.get_key_buffer(layer_id)
v_scratch = pool.get_value_buffer(layer_id)
_check(
rank,
f"mha seq read k l{layer_id}",
k_scratch[rows],
kv_value(all_locs, layer_id, False),
)
_check(
rank,
f"mha seq read v l{layer_id}",
v_scratch[rows],
kv_value(all_locs, layer_id, True),
)
torch.distributed.barrier()
if rank == 0:
print("PASS: MHA page-interleave shard (attn-CP axis)")
@unittest.skipIf(
torch.cuda.device_count() < 2, "page-interleave gather needs 2 CUDA devices"
)
class TestPageInterleaveGatherMultiGpu(CustomTestCase):
"""Real pools, real NCCL, 2 ranks — one mp.spawn per phase.
Separate spawns (and separate ports) because each phase builds its own
process group with a different attention-CP width.
"""
def test_mla_shard_over_attention_tp(self):
mp.spawn(_run_mla, args=(_GATHER_WORLD, 29811), nprocs=_GATHER_WORLD, join=True)
def test_mha_shard_over_attention_cp(self):
mp.spawn(_run_mha, args=(_GATHER_WORLD, 29812), nprocs=_GATHER_WORLD, join=True)
if __name__ == "__main__":
unittest.main()