diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 926f9f2030..d0f0d41787 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -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", diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 13ad6367a5..a9569aadba 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -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) diff --git a/lightllm/server/router/req_queue/dp_balancer/__init__.py b/lightllm/server/router/req_queue/dp_balancer/__init__.py index 34f994f8a2..f17ebc326d 100644 --- a/lightllm/server/router/req_queue/dp_balancer/__init__.py +++ b/lightllm/server/router/req_queue/dp_balancer/__init__.py @@ -2,6 +2,8 @@ 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]): @@ -9,5 +11,10 @@ def get_dp_balancer(args, dp_size_in_node: int, inner_queues: List[BaseQueue]): 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}") diff --git a/lightllm/server/router/req_queue/dp_balancer/cache_aware.py b/lightllm/server/router/req_queue/dp_balancer/cache_aware.py new file mode 100644 index 0000000000..2d27025fc2 --- /dev/null +++ b/lightllm/server/router/req_queue/dp_balancer/cache_aware.py @@ -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() diff --git a/unit_tests/server/router/req_queue/test_dp_cache_aware_balancer.py b/unit_tests/server/router/req_queue/test_dp_cache_aware_balancer.py new file mode 100644 index 0000000000..15453a4d17 --- /dev/null +++ b/unit_tests/server/router/req_queue/test_dp_cache_aware_balancer.py @@ -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")