[Unified Cache][9/N] add opt-in MLA load deduplication for Mooncake Linker (#39565)
Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
co-authored by
Zhangheng
parent
d903351a66
commit
e9300f643e
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user