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
119 changes: 98 additions & 21 deletions integration/hicache/dfkv_hicache.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,68 @@ def _integer(value, field):
raise ValueError(f"{rank_key} must be in [0, {size})")
return size, rank

def _resolve_parallel_coordinates(
cfg: dict,
) -> tuple[tuple[int, int], tuple[int, int], Optional[int]]:
"""Resolve SGLang's physical CP coordinates, with explicit overrides.

Recent SGLang releases do not copy attention-CP/DCP coordinates into the
dynamic HiCache backend's ``extra_config``. Discover them from the initialized
parallel state so sharded ranks cannot alias the default ``0/1`` keys.
"""
parallel = None
try:
try:
from sglang.srt.runtime_context import get_parallel
except ImportError:
# Compatibility with releases that re-export runtime context through
# the distributed package.
from sglang.srt.distributed import get_parallel

parallel = get_parallel()
except (ImportError, AttributeError, AssertionError, RuntimeError):
pass

resolved = {}
runtime_fields = {
"pcp": ("attn_cp_size", "attn_cp_rank"),
"dcp": ("attn_dcp_size", "attn_dcp_rank"),
}
for name, (size_attr, rank_attr) in runtime_fields.items():
if f"{name}_size" in cfg or f"{name}_rank" in cfg:
resolved[name] = _physical_axis(cfg, name)
elif parallel is not None:
resolved[name] = _physical_axis(
{
f"{name}_size": getattr(parallel, size_attr),
f"{name}_rank": getattr(parallel, rank_attr),
},
name,
)
else:
resolved[name] = (1, 0)

attn_tp_rank = (
int(getattr(parallel, "attn_tp_rank"))
if parallel is not None and hasattr(parallel, "attn_tp_rank")
else None
)
return resolved["pcp"], resolved["dcp"], attn_tp_rank


def _is_mla_replica_writer(
is_mla: bool,
tp_rank: int,
attn_tp_rank: Optional[int],
dcp_size: int,
dcp_rank: int,
) -> bool:
"""Elect one writer for each physical PCP/DCP shard."""
if not is_mla:
return True
rank = attn_tp_rank if attn_tp_rank is not None else tp_rank
return rank == (dcp_rank if dcp_size > 1 else 0)


def resolve_node_dedup(cfg_value, env_value, is_mla: bool, tp_size: int):
"""Decide DFKV_CLIENT_NODE_DEDUP: (value_to_set | None, auto_enabled).
Expand Down Expand Up @@ -281,10 +343,24 @@ def __init__(self, storage_config: HiCacheStorageConfig, kwargs: Optional[dict]
self.tp_size = int(storage_config.tp_size)
self.is_mla = bool(storage_config.is_mla_model)
# PCP and DCP split one logical page into distinct physical shards.
# A guessed rank would alias those shards, so multi-rank axes require an
# explicit bounded coordinate before the native client is opened.
self.pcp_size, self.pcp_rank = _physical_axis(cfg, "pcp")
self.dcp_size, self.dcp_rank = _physical_axis(cfg, "dcp")
# SGLang does not forward those coordinates through extra_config, so use
# its initialized parallel state unless the operator explicitly supplied
# a bounded override.
(
(self.pcp_size, self.pcp_rank),
(self.dcp_size, self.dcp_rank),
attn_tp_rank,
) = _resolve_parallel_coordinates(cfg)
# MLA is replicated only inside one effective attention-TP subgroup.
# DCP ranks own distinct interleaved shards, so each DCP coordinate elects
# its own writer; PCP ranks independently repeat the same election.
self._mla_replica_writer = _is_mla_replica_writer(
self.is_mla,
self.tp_rank,
attn_tp_rank,
self.dcp_size,
self.dcp_rank,
)
self._device_direct_requested = _truthy(
os.environ.get("SGLANG_HICACHE_L2_BYPASS"))
self._device_registration_failed = False
Expand Down Expand Up @@ -1025,8 +1101,9 @@ def batch_set_v1(self, keys, host_indices, extra_info=None) -> List[bool]:
n = len(keys)
with _tracing.span("batch_set_v1", n) as _sp, \
access_log("batch_set_v1", lambda: f"{self._alog_tag} {n} keys") as r:
# MLA backup_skip: latent is replicated across TP, only rank 0 writes.
if self.is_mla and self.tp_rank != 0:
# MLA latent is replicated only within one attention-TP subgroup.
# PCP/DCP shards elect an independent writer per physical coordinate.
if self.is_mla and not self._mla_replica_writer:
r.result = "backup_skip"
if _sp:
_sp.attrs = {"dfkv.backup_skip": True}
Expand Down Expand Up @@ -1190,8 +1267,9 @@ def batch_set_v1_device(self, keys, device_indices, extra_info=None) -> List[boo
with _tracing.span("batch_set_v1_device", n) as _sp, \
access_log("batch_set_v1_device",
lambda: f"{self._alog_tag} {n} keys") as r:
# MLA backup_skip: latent is replicated across TP, only rank 0 writes.
if self.is_mla and self.tp_rank != 0:
# MLA latent is replicated only within one attention-TP subgroup.
# PCP/DCP shards elect an independent writer per physical coordinate.
if self.is_mla and not self._mla_replica_writer:
r.result = "backup_skip"
if _sp:
_sp.attrs = {"dfkv.backup_skip": True}
Expand Down Expand Up @@ -1464,9 +1542,12 @@ def _v2_io(self, transfers, putting):
# Replication is inferred from the registered host layout, not a
# model-name allowlist, and _pool_keys carries the matching
# component/rank coordinates. Thus the write gate and collision
# isolation derive from the same physical contract.
if (putting and self.is_mla and self.tp_rank != 0
and self._pool_is_replicated(name)):
if (
putting
and self.is_mla
and not self._mla_replica_writer
and self._pool_is_replicated(name)
):
results[name] = [True] * len(keys)
continue
pool = self.registered_pools[name]
Expand Down Expand Up @@ -1569,8 +1650,7 @@ def _kv_device_set(self, keys, device_indices):
access-log wrapper, so batch_set_v2_device can reuse it for the anchor KV.
Returns (per_page_bools, nbytes, seconds). MLA backup_skip on non-zero TP
rank (replicated latent) short-circuits to all-True, no I/O."""
n = len(keys)
if self.is_mla and self.tp_rank != 0:
if self.is_mla and not self._mla_replica_writer:
return [True] * n, 0, 0.0
from sglang.srt.mem_cache.device_page_meta import (
get_device_page_buffer_meta,
Expand Down Expand Up @@ -1684,8 +1764,7 @@ def _sidecar_device_set(self, name, keys, device_indices):
indexer is layer-first & PAGE-indexed; get_device_sidecar_page_buffer_meta
yields its per-layer page-row segments from the registered sidecar device
pool. Returns (per_page_bools, nbytes, seconds)."""
n = len(keys)
if self.is_mla and self.tp_rank != 0:
if self.is_mla and not self._mla_replica_writer:
return [True] * n, 0, 0.0
from sglang.srt.mem_cache.device_page_meta import (
get_device_sidecar_page_buffer_meta,
Expand Down Expand Up @@ -1754,8 +1833,7 @@ def batch_set_v1_device_draft(self, keys, device_indices, extra_info=None) -> Li
)
seg_ptrs, seg_sizes = get_device_page_buffer_meta(
self.mem_pool_device_draft, device_indices)
sub = len(seg_ptrs) // n if n else 1
if sub == 1 and self.tp_rank != 0:
if sub == 1 and not self._mla_replica_writer:
r.result = "backup_skip"
return [True] * n
stride, sks, sp, ss = self._flatten_device(
Expand Down Expand Up @@ -1829,8 +1907,8 @@ def batch_set_v2_device(
kv_res, kv_bytes, kv_secs = self._kv_device_set(kv_keys, kv_device_indices)
results = {"kv": kv_res}
# Main-KV device write reports on_set (v1-device), matching stock DSA's
# anchor attribution; skip the metric on the MLA rank!=0 no-op.
if n and not (self.is_mla and self.tp_rank != 0):
# anchor attribution; skip the metric on a replicated MLA follower.
if n and (not self.is_mla or self._mla_replica_writer):
self._metrics.on_set(pages=n, ok_pages=sum(kv_res),
nbytes=kv_bytes, seconds=kv_secs)
# Sidecar. Task 4: a DEVICE sidecar transfer (device_indices set,
Expand All @@ -1852,8 +1930,7 @@ def batch_set_v2_device(
results.update(side)
side_pages += sum(len(rs) for rs in side.values())
side_ok += sum(sum(rs) for rs in side.values())
side_bytes += self._v2_io_bytes; side_secs += self._v2_io_seconds
if side_pages and not (self.is_mla and self.tp_rank != 0):
if side_pages and (not self.is_mla or self._mla_replica_writer):
self._metrics.on_set_v2(
pages=side_pages, ok_pages=side_ok,
nbytes=side_bytes, seconds=side_secs)
Expand Down
70 changes: 70 additions & 0 deletions integration/hicache/tests/test_sg_width_namespace.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,24 @@ def _stub_observability() -> None:

_stub_observability()

@contextlib.contextmanager
def _parallel_runtime(**coordinates):
module_name = "sglang.srt.runtime_context"
previous = sys.modules.get(module_name)
module = types.ModuleType(module_name)
parallel = types.SimpleNamespace(**coordinates)
module.get_parallel = lambda: parallel
sys.modules[module_name] = module
try:
yield
finally:
if previous is None:
sys.modules.pop(module_name, None)
else:
sys.modules[module_name] = previous




class _FakeLib:
"""dfkv_max_sg_segs is the only handle _sg_width() needs."""
Expand Down Expand Up @@ -129,6 +147,7 @@ def _mk_instance(width: int, mla: bool = True):
inst.pcp_rank = 0
inst.dcp_size = 1
inst.dcp_rank = 0
inst._mla_replica_writer = False
inst.mem_pool_device = object() # non-None → device (L2-bypass) mode
inst._metrics = _FakeMetrics()
inst._alog_tag = "test"
Expand All @@ -141,6 +160,57 @@ def _mk_instance(width: int, mla: bool = True):
# ---------------------------------------------------------------------------


class TestParallelCoordinates(unittest.TestCase):
def test_discovers_sglang_pcp_dcp_coordinates(self):
with _parallel_runtime(
attn_cp_size=8,
attn_cp_rank=3,
attn_dcp_size=2,
attn_dcp_rank=1,
attn_tp_rank=0,
):
pcp, dcp, attn_tp_rank = H._resolve_parallel_coordinates({})
self.assertEqual(pcp, (8, 3))
self.assertEqual(dcp, (2, 1))
self.assertEqual(attn_tp_rank, 0)

def test_explicit_coordinates_override_runtime_axes(self):
with _parallel_runtime(
attn_cp_size=8,
attn_cp_rank=3,
attn_dcp_size=2,
attn_dcp_rank=1,
attn_tp_rank=2,
):
pcp, dcp, attn_tp_rank = H._resolve_parallel_coordinates(
{
"pcp_size": 4,
"pcp_rank": 1,
"dcp_size": 1,
"dcp_rank": 0,
}
)
self.assertEqual(pcp, (4, 1))
self.assertEqual(dcp, (1, 0))
self.assertEqual(attn_tp_rank, 2)

def test_cp_rank_changes_physical_key(self):
left = _mk_instance(29)
left.pcp_size = 8
left.pcp_rank = 0
right = _mk_instance(29)
right.pcp_size = 8
right.pcp_rank = 1
self.assertNotEqual(left._keys("shared-page"), right._keys("shared-page"))

def test_mla_writer_is_elected_per_physical_cp_shard(self):
self.assertFalse(H._is_mla_replica_writer(True, 3, None, 1, 0))
self.assertTrue(H._is_mla_replica_writer(True, 3, 0, 1, 0))
self.assertTrue(H._is_mla_replica_writer(True, 3, 3, 8, 3))
self.assertFalse(H._is_mla_replica_writer(True, 7, 7, 4, 3))
self.assertTrue(H._is_mla_replica_writer(False, 3, 2, 8, 2))


class TestSgGroupKey(unittest.TestCase):
def test_default_width_is_explicit(self):
inst = _mk_instance(29)
Expand Down
Loading