Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions python/sglang/srt/arg_groups/fields/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,6 +216,10 @@ class Memory(msgspec.Struct):
choices=["mooncake", "mori"],
),
] = "mooncake"
enable_linker_mla_dedup: A[
bool,
"Load replicated MLA KV on rank 0 and broadcast each layer with the Mooncake linker.",
] = False

# -------------------------------------------------------------------------
# Hierarchical sparse attention
Expand Down
5 changes: 5 additions & 0 deletions python/sglang/srt/arg_groups/hicache_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,11 @@ def handle_hicache(server_args: Any):
2) Storage <-> layout compatibility (may rewrite layout).
"""
cfg = resolving_view(server_args)
if cfg.enable_linker_mla_dedup and (
not cfg.enable_unified_cache_external_linker
or cfg.unified_cache_external_linker_backend != "mooncake"
):
raise ValueError("--enable-linker-mla-dedup requires the Mooncake linker.")
if cfg.enable_unified_cache_external_linker:
if cfg.enable_hierarchical_cache:
raise ValueError(
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import hashlib
import logging
import threading
from concurrent.futures import Future
Expand All @@ -19,6 +20,9 @@
from sglang.srt.mem_cache.hybrid_cache.linker_pool_assembler import (
resolve_hybrid_device_pool_group,
)
from sglang.srt.mem_cache.unified_cache.linker_mla_dedup import (
LinkerMLADedupBroadcaster,
)
from sglang.srt.mem_cache.unified_cache.unified_cache_linker import UnifiedCacheLinker
from sglang.srt.runtime_context import (
get_memory,
Expand All @@ -44,8 +48,9 @@ def _storage_suffix(
class LayerWiseLoadCounter:
"""CPU completion counter compatible with KV pools' layer wait hook."""

def __init__(self, num_layers: int):
def __init__(self, num_layers: int, on_layer_ready=None):
self.num_layers = num_layers
self.on_layer_ready = on_layer_ready
self.producer_index = -1
self.consumer_index = -1
self.futures: dict[int, list[Future]] = {}
Expand Down Expand Up @@ -73,6 +78,8 @@ def wait_until(self, threshold: int) -> None:
return
try:
futures[threshold].result()
if self.on_layer_ready is not None:
self.on_layer_ready(index, threshold)
except BaseException as error:
raise RuntimeError("Mooncake layer-wise KV load failed.") from error
finally:
Expand Down Expand Up @@ -158,7 +165,21 @@ def __init__(
)

self.register_buffers()
self.layer_done_counter = LayerWiseLoadCounter(self.num_layers)
self.mla_broadcaster = None
self.tp_group = tp_group
self.broadcast_loads = {}
self.broadcast_events = []
if server_args.enable_linker_mla_dedup and rank_replicated and tp_size > 1:
self.mla_broadcaster = LinkerMLADedupBroadcaster.build(
self.pool_group, params.tp_cache_group, params.attn_tp_cache_group
)
logger.info(
"MLA linker rank-0 loading enabled: tp_rank=%d/%d", tp_rank, tp_size
)
self.layer_done_counter = LayerWiseLoadCounter(
self.num_layers,
self._broadcast_loaded_layers if self.mla_broadcaster else None,
)
if PoolName.MAMBA in self.pools:
params.req_to_token_pool.register_layer_transfer_counter(
self.layer_done_counter
Expand Down Expand Up @@ -242,6 +263,9 @@ def cancel_queued_load(self, rid: str) -> bool:
return False

def num_completed_loads(self) -> int:
while self.broadcast_events and self.broadcast_events[0][0].query():
_, rids = self.broadcast_events.pop(0)
self.completed_loads.put(rids)
return self.completed_loads.qsize()

def pop_completed_load(self) -> list[str]:
Expand All @@ -256,34 +280,95 @@ def freeze_gc_once(self) -> None:
self.gc_frozen = True

def start_layer_wise_loading(self) -> int:
broadcaster = self.mla_broadcaster
if broadcaster is not None:
# Insert can adopt different pages on different ranks. Compare the
# logical load order (not local slots), including empty batches,
# before deciding collectively whether this batch can broadcast.
plan = [
(rid, [(t.name, list(t.keys)) for t in transfers])
for rid, transfers in self.pending_loads.items()
]
digest = hashlib.sha256(repr(plan).encode()).digest()
digests = [None] * torch.distributed.get_world_size(self.tp_group)
torch.distributed.all_gather_object(digests, digest, group=self.tp_group)
if any(other != digest for other in digests):
logger.warning("MLA linker load plans differ; using all-rank reads.")
broadcaster = None
if not self.pending_loads:
return -1
self.freeze_gc_once()
pending = self.pending_loads
self.pending_loads = {}

counter_index = self.layer_done_counter.update_producer()
if broadcaster is not None:
indices = {}
for transfers in pending.values():
for transfer in transfers:
indices.setdefault(transfer.name, []).append(transfer.host_indices)
prepared = broadcaster.prepare_broadcast(
{name: torch.cat(parts) for name, parts in indices.items()},
device_module.current_stream(),
)
if not broadcaster.is_src:
for layer in range(self.num_layers):
self.layer_done_counter.complete(counter_index, layer)
ready_event = device_module.Event()
ready_event.record()
self.load_queue.put((counter_index, pending, ready_event))
if broadcaster is not None:
self.broadcast_loads[counter_index] = (
0,
prepared,
list(pending),
ready_event,
)
if broadcaster is None or broadcaster.is_src:
self.load_queue.put((counter_index, pending, ready_event))
self.stats["load"] += len(pending)
return counter_index

def _broadcast_loaded_layers(self, index: int, threshold: int) -> None:
batch = self.broadcast_loads.get(index)
if batch is None:
return
first, prepared, rids, ready_event = batch
if first == 0:
device_module.current_stream().wait_event(ready_event)
# KV access may wait several times per layer, or skip sparse layers.
# Launch each broadcast exactly once, in order, on the forward stream.
for layer in range(first, threshold + 1):
self.mla_broadcaster.broadcast_loaded_layer(layer, prepared)
if threshold == self.num_layers - 1:
event = device_module.Event()
event.record()
self.broadcast_events.append((event, rids))
del self.broadcast_loads[index]
else:
self.broadcast_loads[index] = (
max(first, threshold + 1),
prepared,
rids,
ready_event,
)

def load_thread_func(self) -> None:
while True:
task = self.load_queue.get()
try:
if task is None:
return
counter_index, pending, ready_event = task
replicated = counter_index in self.broadcast_loads
try:
ready_event.synchronize()
self.load_layer_wise(counter_index, list(pending.values()))
except BaseException as error:
self.layer_done_counter.fail(counter_index, error)
logger.exception("Mooncake layer-wise load batch failed")
finally:
self.completed_loads.put(list(pending))
if not replicated:
self.completed_loads.put(list(pending))
finally:
self.load_queue.task_done()

Expand Down Expand Up @@ -397,6 +482,10 @@ def reset(self) -> None:
self.pending_loads.clear()
self.load_queue.join()
self.offload_queue.join()
for event, _ in self.broadcast_events:
event.synchronize()
self.broadcast_events.clear()
self.broadcast_loads.clear()
while True:
try:
self.offload_results.get_nowait()
Expand All @@ -417,3 +506,5 @@ def close(self) -> None:
self.offload_thread.join()
logger.info("Mooncake direct linker stats: %s", self.stats)
self.storage.close()
if self.mla_broadcaster is not None:
self.mla_broadcaster.destroy()
85 changes: 85 additions & 0 deletions python/sglang/srt/mem_cache/unified_cache/linker_mla_dedup.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
"""Backend-independent MLA broadcast adapter for external cache linkers."""

from __future__ import annotations

from types import SimpleNamespace

import torch

from sglang.srt.mem_cache.mla_host_dedup import MLAHostDedupBroadcaster
from sglang.srt.utils import get_device_module

device_module = get_device_module()


class LinkerMLADedupBroadcaster(MLAHostDedupBroadcaster):
"""Adapt hybrid pool geometry to HiCache's existing broadcast primitive.

Native MLA/DSA pools can use MLAHostDedupBroadcaster directly. This adapter
covers multiple physical pools and sparse layers (e.g. DSV4). The caller
must provide the same pools and logical page order on all ranks; only
physical slots may differ. Call on the forward stream after source KV is
loaded. Read completion, failure coordination and scheduling stay outside.
"""

def __init__(self, pool_group, group, src_global_rank):
if not pool_group.rank_replicated:
raise ValueError("Linker broadcasts require rank-replicated pools.")
self.pools = pool_group.entry_map
self.buffers = {
name: [
[buf.view(torch.uint8).view(buf.shape[0], -1) for buf in component]
for component in pool.components
]
for name, pool in self.pools.items()
}
first = next(iter(self.buffers.values()))[0][0]
# Reuse HiCache's allocation and token-based chunk budget, even when
# the physical buffers store page rows rather than individual tokens.
bytes_per_token = max(
(size + pool.page_size - 1) // pool.page_size
for pool in self.pools.values()
for component in pool.buffer_meta
for _, _, size in component
)
super().__init__(
SimpleNamespace(
device=first.device,
layer_num=pool_group.num_layers,
kv_cache_dim=bytes_per_token,
kv_buffer=[first],
),
group,
src_global_rank,
)
# A physical page cannot be split into smaller rows by _bcast_layer.
max_row_bytes = max(
buf.shape[1]
for components in self.buffers.values()
for component in components
for buf in component
)
self.kv_staging.resize_(max(self.kv_staging.numel(), max_row_bytes))

def prepare_broadcast(self, indices_by_pool, load_stream):
"""Map already-translated device slots to each physical pool's rows."""
prepared = {}
for name, indices in indices_by_pool.items():
pool = self.pools[name]
rows = torch.tensor(pool.prepare_locations(indices), dtype=torch.int64)
rows = (rows[:, None] + torch.arange(pool._row_span)).flatten()
prepared[name] = super().prepare_broadcast(rows, load_stream)
return prepared

def broadcast_loaded_layer(self, layer_id, prepared):
for name, pool in self.pools.items():
layer = pool.layer_mapping.get(layer_id)
if layer is None or name not in prepared:
continue
indices, _ = prepared[name]
if indices.is_cuda:
indices.record_stream(device_module.current_stream())
for buffers in self.buffers[name]:
self._bcast_layer(
buffers, self.kv_staging, indices, buffers[layer].shape[1], layer
)
19 changes: 19 additions & 0 deletions test/registered/unit/server_args/test_server_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -1901,6 +1901,25 @@ def test_enable_ssl_refresh_with_ssl_accepted(self, _mock_isfile):


class TestHiCacheArgs(unittest.TestCase):
def test_linker_mla_dedup_requires_mooncake_linker(self):
for enabled, linker, backend in (
(False, False, "mooncake"),
(True, True, "mooncake"),
(True, False, "mooncake"),
(True, True, "mori"),
):
with self.subTest(enabled=enabled, linker=linker, backend=backend):
args = self._make_args(
enable_linker_mla_dedup=enabled,
enable_unified_cache_external_linker=linker,
unified_cache_external_linker_backend=backend,
)
if enabled and (not linker or backend != "mooncake"):
with self.assertRaisesRegex(ValueError, "requires the Mooncake"):
handle_hicache(args)
else:
handle_hicache(args)

def _make_args(self, **overrides) -> ServerArgs:
# Not resolved: a dummy model path takes the pipeline's early return,
# so `_handle_hicache` would never run. Its one prerequisite (the
Expand Down
Loading