diff --git a/python/sglang/srt/arg_groups/fields/memory.py b/python/sglang/srt/arg_groups/fields/memory.py index 29c2b49dd..be8d204e6 100644 --- a/python/sglang/srt/arg_groups/fields/memory.py +++ b/python/sglang/srt/arg_groups/fields/memory.py @@ -220,6 +220,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 diff --git a/python/sglang/srt/arg_groups/hicache_hook.py b/python/sglang/srt/arg_groups/hicache_hook.py index f9e4df805..757eaa92e 100644 --- a/python/sglang/srt/arg_groups/hicache_hook.py +++ b/python/sglang/srt/arg_groups/hicache_hook.py @@ -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( diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_direct_linker.py b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_direct_linker.py index 1c4d573a7..09d6b45bf 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_direct_linker.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_direct_linker.py @@ -1,5 +1,6 @@ from __future__ import annotations +import hashlib import logging import threading from concurrent.futures import Future @@ -19,6 +20,9 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( 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, @@ -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]] = {} @@ -73,6 +78,8 @@ class LayerWiseLoadCounter: 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: @@ -158,7 +165,21 @@ class MooncakeDirectLinker(UnifiedCacheLinker): ) 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 @@ -242,6 +263,9 @@ class MooncakeDirectLinker(UnifiedCacheLinker): 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]: @@ -256,6 +280,21 @@ class MooncakeDirectLinker(UnifiedCacheLinker): 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() @@ -263,12 +302,56 @@ class MooncakeDirectLinker(UnifiedCacheLinker): 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() @@ -276,6 +359,7 @@ class MooncakeDirectLinker(UnifiedCacheLinker): 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())) @@ -283,7 +367,8 @@ class MooncakeDirectLinker(UnifiedCacheLinker): 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() @@ -397,6 +482,10 @@ class MooncakeDirectLinker(UnifiedCacheLinker): 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() @@ -417,3 +506,5 @@ class MooncakeDirectLinker(UnifiedCacheLinker): 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() diff --git a/python/sglang/srt/mem_cache/unified_cache/linker_mla_dedup.py b/python/sglang/srt/mem_cache/unified_cache/linker_mla_dedup.py new file mode 100644 index 000000000..2c90f20d3 --- /dev/null +++ b/python/sglang/srt/mem_cache/unified_cache/linker_mla_dedup.py @@ -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 + ) diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 5972c2529..e82fbddc0 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -1995,6 +1995,25 @@ class TestSSLArgs(unittest.TestCase): 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