Skip to content
Open
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: 2 additions & 2 deletions lightllm/server/api_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -286,8 +286,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser:
"--dp_balancer",
type=str,
default="bs_balancer",
choices=["round_robin", "bs_balancer"],
help="the dp balancer type, default is bs_balancer",
choices=["round_robin", "bs_balancer", "cache_aware"],
help="the DP balancer type; cache_aware adds token-prefix affinity, default is bs_balancer",
)
parser.add_argument(
"--max_req_total_len",
Expand Down
2 changes: 1 addition & 1 deletion lightllm/server/core/objs/start_args_type.py
Original file line number Diff line number Diff line change
Expand Up @@ -224,7 +224,7 @@ class StartArgs:
multinode_httpmanager_port: int = field(default=12345)

disable_shm_warning: bool = field(default=False)
dp_balancer: str = field(default="bs_balancer", metadata={"choices": ["round_robin", "bs_balancer"]})
dp_balancer: str = field(default="bs_balancer", metadata={"choices": ["round_robin", "bs_balancer", "cache_aware"]})
enable_fused_shared_experts: bool = field(default=False)
enable_mps: bool = field(default=False)
multinode_router_gloo_port: int = field(default=20001)
Expand Down
7 changes: 7 additions & 0 deletions lightllm/server/router/req_queue/dp_balancer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,19 @@
from typing import List
from lightllm.server.router.req_queue.base_queue import BaseQueue
from .bs import DpBsBalancer
from .cache_aware import DpCacheAwareBalancer, DpCacheAwareConfig
from lightllm.utils.config_utils import is_hybrid_att_model


def get_dp_balancer(args, dp_size_in_node: int, inner_queues: List[BaseQueue]):
if args.dp_balancer == "round_robin":
return RoundRobinDpBalancer(dp_size_in_node, inner_queues)
elif args.dp_balancer == "bs_balancer":
return DpBsBalancer(dp_size_in_node, inner_queues)
elif args.dp_balancer == "cache_aware":
if args.disable_dynamic_prompt_cache:
raise ValueError("cache_aware DP balancing requires dynamic prompt cache")
block_size = args.linear_att_hash_page_size if is_hybrid_att_model(args.model_dir) else args.page_size
return DpCacheAwareBalancer(dp_size_in_node, inner_queues, DpCacheAwareConfig(block_size=block_size))
else:
raise ValueError(f"Invalid dp balancer: {args.dp_balancer}")
178 changes: 178 additions & 0 deletions lightllm/server/router/req_queue/dp_balancer/cache_aware.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,178 @@
"""DP-local cache-affinity routing based on bounded token-prefix history.

The router owns this heuristic index. It records dispatch history rather than querying
the infer processes' radix trees, so stale entries can only affect placement, not KV
cache correctness.
"""

from __future__ import annotations

import random
from collections import OrderedDict
from dataclasses import dataclass
from typing import List, Optional, Tuple

import xxhash

from lightllm.server.router.batch import Batch, Req
from lightllm.server.router.req_queue.base_queue import BaseQueue

from .base import DpBalancer


PrefixHash = Tuple[int, int]


@dataclass(slots=True)
class DpCacheAwareConfig:
block_size: int
cache_threshold: float = 0.5
balance_rel_threshold: float = 1.8
max_cache_entries: int = 1_000_000
evict_entries: int = 10_000


class TokenPrefixCache:
"""Bounded LRU mapping from cumulative token-prefix hashes to local DP indexes."""

def __init__(self, block_size: int, max_entries: int, evict_entries: int) -> None:
if block_size < 1:
raise ValueError(f"block_size must be >= 1, got {block_size}")
if max_entries < 0:
raise ValueError(f"max_entries must be >= 0, got {max_entries}")
if evict_entries < 1:
raise ValueError(f"evict_entries must be >= 1, got {evict_entries}")
self.block_size = block_size
self.max_entries = max_entries
self.evict_entries = evict_entries
self._prefix_to_dp: OrderedDict[int, int] = OrderedDict()

def hash_prefixes(self, prompt_ids) -> List[PrefixHash]:
cacheable_token_count = max(0, len(prompt_ids) - 1)
cacheable_token_count = cacheable_token_count // self.block_size * self.block_size
if cacheable_token_count == 0:
return []

token_view = memoryview(prompt_ids)
item_size = token_view.itemsize
prompt_bytes = token_view.cast("B")
# A collision only changes a routing hint; it cannot affect cache correctness.
# xxh3-64 keeps the 1M-entry index compact and hashes 1M-token prompts faster.
hasher = xxhash.xxh3_64()
prefix_hashes = []
for start in range(0, cacheable_token_count, self.block_size):
end = start + self.block_size
hasher.update(prompt_bytes[start * item_size : end * item_size])
prefix_hashes.append((hasher.intdigest(), end))
return prefix_hashes

def match(self, prefix_hashes: List[PrefixHash]) -> Tuple[Optional[int], int]:
for prefix_hash, token_count in reversed(prefix_hashes):
try:
dp_index = self._prefix_to_dp[prefix_hash]
except KeyError:
continue
self._prefix_to_dp.move_to_end(prefix_hash)
return dp_index, token_count
return None, 0

def insert(self, prefix_hashes: List[PrefixHash], dp_index: int, start_index: int = 0) -> None:
for prefix_index in range(start_index, len(prefix_hashes)):
prefix_hash = prefix_hashes[prefix_index][0]
self._prefix_to_dp[prefix_hash] = dp_index
self._prefix_to_dp.move_to_end(prefix_hash)

if len(self._prefix_to_dp) > self.max_entries:
evict_count = len(self._prefix_to_dp) - self.max_entries + self.evict_entries
for _ in range(min(evict_count, len(self._prefix_to_dp))):
self._prefix_to_dp.popitem(last=False)

def __len__(self) -> int:
return len(self._prefix_to_dp)


class DpCacheAwareBalancer(DpBalancer):
"""Route matching token prefixes to the same local DP unless load requires rebalancing."""

def __init__(
self,
dp_size_in_node: int,
inner_queues: List[BaseQueue],
config: DpCacheAwareConfig,
) -> None:
super().__init__(dp_size_in_node, inner_queues)
self.config = config
self.prefix_cache = TokenPrefixCache(
block_size=self.config.block_size,
max_entries=self.config.max_cache_entries,
evict_entries=self.config.evict_entries,
)

def assign_reqs_to_dp(self, current_batch: Batch, reqs_waiting_for_dp_index: List[List[Req]]) -> None:
if not reqs_waiting_for_dp_index:
return

current_load_per_dp = [0 for _ in range(self.dp_size_in_node)]
if current_batch is not None:
current_load_per_dp = current_batch.get_all_dp_req_num()
total_load_per_dp = [
current_load_per_dp[dp_index] + len(self.inner_queues[dp_index].waiting_req_list)
for dp_index in range(self.dp_size_in_node)
]

for req_group in reqs_waiting_for_dp_index:
first_req = req_group[0]
prefix_hashes = []
if not first_req.sample_params.disable_prompt_cache:
linked_prompt_ids = False
if not hasattr(first_req, "shm_prompt_ids"):
first_req.link_prompt_ids_shm_array()
linked_prompt_ids = True
try:
prefix_hashes = self.prefix_cache.hash_prefixes(first_req.get_prompt_ids_numpy())
finally:
if linked_prompt_ids:
first_req.shm_prompt_ids.detach_shm()
del first_req.shm_prompt_ids

cache_dp_index = None
matched_token_count = 0
if not first_req.sample_params.disable_prompt_cache:
matched_dp_index, matched_token_count = self.prefix_cache.match(prefix_hashes)
match_rate = matched_token_count / first_req.input_len if first_req.input_len else 0.0
if match_rate > self.config.cache_threshold:
cache_dp_index = matched_dp_index

idle_dp_indexes = [dp_index for dp_index, load in enumerate(total_load_per_dp) if load == 0]
if idle_dp_indexes:
if cache_dp_index in idle_dp_indexes:
selected_dp_index = cache_dp_index
else:
selected_dp_index = random.choice(idle_dp_indexes)
else:
min_load = min(total_load_per_dp)
least_loaded_dp_indexes = [
dp_index for dp_index, load in enumerate(total_load_per_dp) if load == min_load
]
least_loaded_dp_index = random.choice(least_loaded_dp_indexes)
if cache_dp_index is None:
selected_dp_index = least_loaded_dp_index
else:
group_load = len(req_group)
cache_projected_load = total_load_per_dp[cache_dp_index] + group_load
least_projected_load = total_load_per_dp[least_loaded_dp_index] + group_load
if cache_projected_load > least_projected_load * self.config.balance_rel_threshold:
selected_dp_index = least_loaded_dp_index
else:
selected_dp_index = cache_dp_index

for req in req_group:
req.sample_params.suggested_dp_index = selected_dp_index
self.inner_queues[selected_dp_index].extend(req_group)
total_load_per_dp[selected_dp_index] += len(req_group)
insert_start_index = 0
if cache_dp_index == selected_dp_index:
insert_start_index = (matched_token_count + self.config.block_size - 1) // self.config.block_size
self.prefix_cache.insert(prefix_hashes, selected_dp_index, start_index=insert_start_index)

reqs_waiting_for_dp_index.clear()
111 changes: 111 additions & 0 deletions unit_tests/server/router/req_queue/test_dp_cache_aware_balancer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,111 @@
from types import SimpleNamespace
from unittest.mock import Mock

import numpy as np
import pytest

from lightllm.server.api_cli import make_argument_parser
from lightllm.server.router.req_queue import dp_balancer
from lightllm.server.router.req_queue.dp_balancer.cache_aware import (
DpCacheAwareBalancer,
DpCacheAwareConfig,
TokenPrefixCache,
)


class Queue:
def __init__(self):
self.waiting_req_list = []

def extend(self, reqs):
self.waiting_req_list.extend(reqs)


def req(tokens, disabled=False):
tokens = np.array(tokens, dtype=np.int64)
return SimpleNamespace(
input_len=len(tokens),
sample_params=SimpleNamespace(disable_prompt_cache=disabled),
shm_prompt_ids=object(),
get_prompt_ids_numpy=lambda: tokens,
)


@pytest.mark.parametrize("hybrid, expected", [(False, 16), (True, 512)])
def test_factory_uses_existing_model_cache_page_size(monkeypatch, hybrid, expected):
monkeypatch.setattr(dp_balancer, "is_hybrid_att_model", lambda _path: hybrid)
args = SimpleNamespace(
dp_balancer="cache_aware",
disable_dynamic_prompt_cache=False,
model_dir="model",
page_size=16,
linear_att_hash_page_size=512,
)
balancer = dp_balancer.get_dp_balancer(args, 2, [Queue(), Queue()])
assert balancer.config.block_size == expected
args.disable_dynamic_prompt_cache = True
with pytest.raises(ValueError, match="requires dynamic prompt cache"):
dp_balancer.get_dp_balancer(args, 2, [Queue(), Queue()])


def test_cli_preserves_default_and_accepts_opt_in_strategy():
parser = make_argument_parser()
assert parser.parse_args([]).dp_balancer == "bs_balancer"
assert parser.parse_args(["--dp_balancer", "cache_aware"]).dp_balancer == "cache_aware"


def test_prefix_hashes_exclude_last_token_and_match_longest_aligned_prefix():
cache = TokenPrefixCache(4, 10, 1)
assert not cache.hash_prefixes(np.arange(4, dtype=np.int64))
first = cache.hash_prefixes(np.arange(9, dtype=np.int64))
cache.insert(first, 1)
assert cache.match(first) == (1, 8)
changed = np.arange(9, dtype=np.int64)
changed[5] = 100
assert cache.match(cache.hash_prefixes(changed)) == (1, 4)


def test_prefix_index_is_bounded_and_evicts_least_recently_used():
cache = TokenPrefixCache(1, 2, 1)
cache.insert([(1, 1), (2, 2)], 0)
assert cache.match([(1, 1)]) == (0, 1)
cache.insert([(3, 1)], 1)
assert len(cache) <= 2
assert cache.match([(2, 2)]) == (None, 0)
assert cache.match([(3, 1)]) == (1, 1)


def test_affinity_preserves_request_group_and_yields_to_overload():
queues = [Queue(), Queue()]
balancer = DpCacheAwareBalancer(2, queues, DpCacheAwareConfig(block_size=4))
tokens = list(range(9))
balancer.prefix_cache.insert(balancer.prefix_cache.hash_prefixes(np.array(tokens, dtype=np.int64)), 1)
group = [req(tokens), req(tokens)]
pending = [group]
balancer.assign_reqs_to_dp(SimpleNamespace(get_all_dp_req_num=lambda: [1, 1]), pending)
assert not pending and queues[1].waiting_req_list == group
assert all(item.sample_params.suggested_dp_index == 1 for item in group)
overloaded = req(tokens)
balancer.assign_reqs_to_dp(SimpleNamespace(get_all_dp_req_num=lambda: [1, 20]), [[overloaded]])
assert overloaded.sample_params.suggested_dp_index == 0


def test_disabled_cache_does_not_attach_prompt_or_publish_affinity():
queues = [Queue(), Queue()]
balancer = DpCacheAwareBalancer(2, queues, DpCacheAwareConfig(block_size=4))
request = req(range(9), disabled=True)
del request.shm_prompt_ids
request.link_prompt_ids_shm_array = Mock(side_effect=AssertionError("must not attach"))
balancer.assign_reqs_to_dp(None, [[request]])
assert len(balancer.prefix_cache) == 0


def test_router_releases_only_prompt_mapping_it_attached():
balancer = DpCacheAwareBalancer(1, [Queue()], DpCacheAwareConfig(block_size=4))
request = req(range(9))
del request.shm_prompt_ids
mapping = SimpleNamespace(detach_shm=Mock())
request.link_prompt_ids_shm_array = lambda: setattr(request, "shm_prompt_ids", mapping)
balancer.assign_reqs_to_dp(None, [[request]])
mapping.detach_shm.assert_called_once_with()
assert not hasattr(request, "shm_prompt_ids")
Loading