[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:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user