[Feat][GLM5.2] Add DSA Cache Layer Split under Prefill CP (#29421)

Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com>
This commit is contained in:
Shijin Zhang
2026-07-09 03:03:56 -07:00
committed by GitHub
parent 336b64ecce
commit 8e54517f02
21 changed files with 1507 additions and 72 deletions
@@ -0,0 +1,83 @@
"""End-to-end GSM8K accuracy test for DSA cache layer split (GLM-5.2).
Layer split shards the DSA GPU KV/indexer cache layers across prefill CP ranks
(``--enable-dsa-cache-layer-split``); non-owner ranks read a layer via an
owner-broadcast into a small remote scratch buffer. It only applies to PD
prefill workers running DSA prefill-CP (a unified server would decode on the
same worker, where non-owner ranks lack the full cache), so this test drives a
PD-disaggregated GLM-5.2 deployment: a layer-split prefill worker running
interleave prefill-CP + layer split, and an ordinary decode worker that receives
full cache shards via PD transfer.
Sized for the 4-GPU B200 runner (prefill TP=2 + decode TP=2) rather than an
8-GPU deployment, since the 8-gpu-b200 runner is nightly-only.
"""
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase,
)
register_cuda_ci(
est_time=1200,
stage="extra-b",
runner_config="4-gpu-b200",
disabled="Temporarily disabled",
)
class TestGLM52DSACacheLayerSplit(PDDisaggregationServerBase, GSM8KMixin):
model = "nvidia/GLM-5.2-NVFP4"
# Full GSM8K test set (1319 questions) with a tight accuracy floor.
gsm8k_accuracy_thres = 0.935
gsm8k_num_questions = 1319
gsm8k_num_threads = 200
gsm8k_num_shots = 0
# Prefill worker: interleave prefill-CP + DSA cache layer split on 2 GPUs
# (TP=2 -> attn_cp_size=2, so KV/indexer layers shard 2-way across CP ranks).
extra_prefill_args = [
"--tp",
"2",
"--dsa-prefill-backend",
"trtllm",
"--kv-cache-dtype",
"fp8_e4m3",
"--enable-dsa-cache-layer-split",
"--enable-prefill-cp",
"--cp-strategy",
"interleave",
"--mem-fraction-static",
"0.85",
"--chunked-prefill-size",
"4096",
"--max-prefill-tokens",
"4096",
]
# Decode worker: ordinary local decode cache on the other 2 GPUs, receives
# full shards via PD transfer.
extra_decode_args = [
"--tp",
"2",
"--dsa-decode-backend",
"trtllm",
"--kv-cache-dtype",
"fp8_e4m3",
"--mem-fraction-static",
"0.85",
"--base-gpu-id",
"2",
]
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.launch_all()
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,123 @@
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.layers.cp.utils import get_layer_owner, get_layer_shard_range
from sglang.srt.mem_cache.dsa_cache_layer_split import LayerSplitDSATokenToKVPool
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestDSALayerShardUtils(CustomTestCase):
def test_balanced_layer_ranges_cover_all_layers_once(self):
ranges = [get_layer_shard_range(rank, 4, 10) for rank in range(4)]
self.assertEqual(ranges, [(0, 3), (3, 6), (6, 8), (8, 10)])
covered = [layer_id for start, end in ranges for layer_id in range(start, end)]
self.assertEqual(covered, list(range(10)))
def test_owner_matches_uneven_layer_ranges(self):
self.assertEqual(
[get_layer_owner(i, 4, 10) for i in range(10)],
[0, 0, 0, 1, 1, 1, 2, 2, 3, 3],
)
def test_empty_tail_shards_have_empty_ranges(self):
ranges = [get_layer_shard_range(rank, 4, 2) for rank in range(4)]
self.assertEqual(ranges, [(0, 1), (1, 2), (2, 2), (2, 2)])
def test_prefetch_uses_sync_fallback_without_dedicated_communicator(self):
broadcasts = []
counter = SimpleNamespace(wait_until=lambda _: self.fail("unexpected wait"))
pool = SimpleNamespace(
remote_kv_layer_id=None,
pending_remote_kv_broadcast=False,
pending_remote_kv_layer_id=None,
layer_broadcast_comm=None,
remote_kv_buffer=object(),
kv_buffer=[object()],
start_layer=0,
_local_layer_idx=lambda layer_id: layer_id,
_is_layer_owned=lambda _: True,
)
def broadcast(tensor, layer_id, *, src_tensor, use_layer_broadcast_comm):
broadcasts.append((tensor, layer_id, src_tensor, use_layer_broadcast_comm))
pool._broadcast_tensor_from_owner = broadcast
# Bind the real method against a lightweight stand-in so the sync
# (no dedicated NCCL comm) fallback path can be exercised on CPU.
LayerSplitDSATokenToKVPool.prefetch_kv_buffer(
pool,
layer_id=0,
layer_transfer_counter=counter,
layer_transfer_idx=3,
)
self.assertEqual(len(broadcasts), 1)
self.assertEqual(pool.remote_kv_layer_id, 0)
def test_finalize_pending_broadcast_promotes_layer_id(self):
# After an async prefetch, finalizing must promote pending -> remote so a
# subsequent read of the same layer reuses the broadcast result.
pool = SimpleNamespace(
pending_remote_kv_broadcast=True,
pending_remote_kv_layer_id=7,
remote_kv_layer_id=None,
device_module=SimpleNamespace(
current_stream=lambda: SimpleNamespace(wait_stream=lambda _stream: None)
),
kv_broadcast_stream=object(),
)
LayerSplitDSATokenToKVPool._finalize_pending_kv_broadcast(
pool, set_remote_layer_id=True
)
self.assertFalse(pool.pending_remote_kv_broadcast)
self.assertEqual(pool.remote_kv_layer_id, 7)
self.assertIsNone(pool.pending_remote_kv_layer_id)
def test_get_broadcastable_kv_buffer_returns_owner_contents(self):
# A non-owner read must return the *owner's* KV bytes, copied into the
# remote scratch buffer by the broadcast. This checks prefetch_kv_buffer
# + _get_broadcastable_kv_buffer surface the correct contents.
layer_num = 4
shard_size = 2
owner_kv = {
layer_id: torch.full((3, 1, 8), float(layer_id + 1))
for layer_id in range(layer_num)
}
remote = torch.zeros((3, 1, 8))
pool = SimpleNamespace(
layer_num=layer_num,
layer_shard_size=shard_size,
start_layer=0,
remote_kv_layer_id=None,
pending_remote_kv_broadcast=False,
pending_remote_kv_layer_id=None,
remote_kv_buffer=remote,
)
pool._local_layer_idx = lambda layer_id: layer_id - pool.start_layer
pool._is_layer_owned = lambda layer_id: True
# kv_buffer holds this rank's owned layers; broadcast copies owner->remote.
pool.kv_buffer = [owner_kv[i] for i in range(layer_num)]
def broadcast(tensor, layer_id, *, src_tensor, use_layer_broadcast_comm=False):
# Simulate the owner writing its layer into the remote scratch buffer.
tensor.copy_(owner_kv[layer_id])
pool._broadcast_tensor_from_owner = broadcast
for layer_id in range(layer_num):
buf = LayerSplitDSATokenToKVPool._get_broadcastable_kv_buffer(
pool, layer_id
)
self.assertTrue(torch.equal(buf, owner_kv[layer_id]))
self.assertEqual(pool.remote_kv_layer_id, layer_id)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,149 @@
"""Multi-GPU integration test for LayerSplitDSATokenToKVPool owner-broadcast.
Spawns ``world`` processes forming a single attention-CP group, builds a tiny
``LayerSplitDSATokenToKVPool`` on each rank, writes a rank-distinct value into
every owned layer, then verifies that reading ANY layer (owned or not) returns
the *owning* rank's bytes -- i.e. the owner-broadcast in
``_get_broadcastable_kv_buffer`` / ``prefetch_kv_buffer`` surfaces correct
contents. Also exercises the DSA indexer broadcast and the async prefetch path.
Registered as a base-c 4-gpu-b200 unit test; uses up to 4 GPUs and skips when
fewer than 2 are visible. Run directly on 2+ GPUs:
CUDA_VISIBLE_DEVICES=0,1 python -m pytest \
test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py
"""
import os
import unittest
import torch
import torch.multiprocessing as mp
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=120, stage="base-c", runner_config="4-gpu-b200")
LAYER_NUM = 4
PAGE_SIZE = 64
KV_LORA_RANK = 512
QK_ROPE = 64
INDEX_HEAD_DIM = 128
SIZE = PAGE_SIZE * 3 # a few pages
PORT = 29711
def _run(rank: int, world: int, port: int):
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)
from sglang.srt.distributed.parallel_state import (
init_distributed_environment,
initialize_model_parallel,
)
from sglang.srt.layers.dp_attention import (
get_attention_cp_rank,
get_attention_cp_size,
)
init_distributed_environment(
world_size=world,
rank=rank,
local_rank=rank,
distributed_init_method=f"tcp://127.0.0.1:{port}",
backend="nccl",
)
initialize_model_parallel(
tensor_model_parallel_size=world,
attention_context_model_parallel_size=world,
)
from sglang.srt.mem_cache.dsa_cache_layer_split import (
LayerSplitDSATokenToKVPool,
)
cp_rank = get_attention_cp_rank()
cp_size = get_attention_cp_size()
assert cp_size == world
pool = LayerSplitDSATokenToKVPool(
SIZE,
page_size=PAGE_SIZE,
kv_lora_rank=KV_LORA_RANK,
dtype=torch.bfloat16,
qk_rope_head_dim=QK_ROPE,
layer_num=LAYER_NUM,
device=f"cuda:{rank}",
index_head_dim=INDEX_HEAD_DIM,
enable_memory_saver=False,
kv_cache_dim=KV_LORA_RANK + QK_ROPE,
layer_shard_rank=cp_rank,
layer_shard_size=cp_size,
)
# Owner writes a layer-distinct constant into each owned kv_buffer layer.
for layer_id in range(LAYER_NUM):
if pool._is_layer_owned(layer_id):
pool.kv_buffer[layer_id].fill_(float(layer_id + 1))
torch.cuda.synchronize()
torch.distributed.barrier()
# Every rank reads every layer; broadcast must surface the owner's value.
ok = True
for layer_id in range(LAYER_NUM):
buf = pool._get_broadcastable_kv_buffer(layer_id)
expected = float(layer_id + 1)
got = buf.float().mean().item()
if abs(got - expected) > 1e-3:
print(f"[rank {rank}] layer {layer_id}: expected {expected}, got {got}")
ok = False
assert ok, f"rank {rank} read stale/incorrect broadcast contents"
# Indexer buffer owner-broadcast: owner writes a layer-distinct value, then
# every rank must read it back for every layer.
for layer_id in range(LAYER_NUM):
if pool._is_layer_owned(layer_id):
pool.index_k_with_scale_buffer[layer_id].fill_(layer_id + 10)
torch.cuda.synchronize()
torch.distributed.barrier()
for layer_id in range(LAYER_NUM):
# invalidate any cached remote copy so the read forces a fresh broadcast
pool.invalidate_index_buffer_for_layer(layer_id)
buf = pool._get_broadcastable_index_buffer(layer_id)
expected = layer_id + 10
got = buf.float().mean().item()
if abs(got - expected) > 1e-3:
print(f"[rank {rank}] index layer {layer_id}: exp {expected}, got {got}")
ok = False
assert ok, f"rank {rank} read stale/incorrect index broadcast contents"
# Async prefetch path: prefetch layer, then read must return owner value.
for layer_id in range(LAYER_NUM):
pool.remote_kv_layer_id = None # force a fresh broadcast
pool.prefetch_kv_buffer(layer_id)
buf = pool._get_broadcastable_kv_buffer(layer_id)
got = buf.float().mean().item()
if abs(got - float(layer_id + 1)) > 1e-3:
print(f"[rank {rank}] prefetch layer {layer_id}: got {got}")
ok = False
assert ok, f"rank {rank} prefetch path returned incorrect contents"
print(f"[rank {rank}] OK: all {LAYER_NUM} layers read correct owner contents")
torch.distributed.barrier()
class TestLayerSplitDSABroadcast(CustomTestCase):
def test_owner_broadcast(self):
world = min(4, torch.cuda.device_count())
if world < 2:
self.skipTest("LayerSplitDSATokenToKVPool broadcast test needs >= 2 GPUs")
mp.spawn(_run, args=(world, PORT), nprocs=world, join=True)
if __name__ == "__main__":
unittest.main()
@@ -49,6 +49,15 @@ def _ptr_key_from_tensor(ptrs: torch.Tensor) -> tuple[int, ...]:
return tuple(int(ptr) for ptr in ptrs.cpu().tolist())
def _device_pool_stub(*, layer_num: int, **fields) -> SimpleNamespace:
"""Minimal device-pool stand-in with layer-split fields real pools expose."""
return SimpleNamespace(
layer_num=layer_num,
layer_shard_enabled=False,
**fields,
)
def _cpu_staged_lf_pf_copy(
src_registry,
*,
@@ -192,7 +201,8 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
]
expected_k = [layer[device_indices].clone() for layer in k_layers]
expected_v = [layer[device_indices].clone() for layer in v_layers]
device_pool = SimpleNamespace(
device_pool = _device_pool_stub(
layer_num=layer_num,
k_buffer=k_layers,
v_buffer=v_layers,
k_data_ptrs=torch.tensor(
@@ -293,7 +303,8 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
for layer_id in range(layer_num)
]
expected = [layer[device_indices].clone() for layer in device_layers]
device_pool = SimpleNamespace(
device_pool = _device_pool_stub(
layer_num=layer_num,
kv_buffer=device_layers,
data_ptrs=torch.tensor(
[layer.data_ptr() for layer in device_layers], dtype=torch.uint64
@@ -301,6 +312,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
)
host = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost)
host.device_pool = device_pool
host.layout = "page_first"
host.page_size = 1
host.layer_num = layer_num
@@ -582,9 +594,13 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
for layer_id in range(layer_num)
]
expected = [buffer[device_page_indices].clone() for buffer in device_layers]
device_pool = SimpleNamespace(index_k_with_scale_buffer=device_layers)
device_pool = _device_pool_stub(
layer_num=layer_num,
index_k_with_scale_buffer=device_layers,
)
host = DSAIndexerPoolHost.__new__(DSAIndexerPoolHost)
host.device_pool = device_pool
host.layout = "page_first"
host.page_size = page_size
host.layer_num = layer_num
@@ -114,6 +114,7 @@ def _make_model_runner(
sa.disaggregation_mode = disaggregation_mode
sa.max_running_requests = max_running_requests
sa.disaggregation_decode_extra_slots = disaggregation_decode_extra_slots
sa.enable_dsa_cache_layer_split = False
mr.server_args = sa
spec = MagicMock()