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
14 changes: 8 additions & 6 deletions lightllm/server/httpserver_for_pd_master/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,9 +211,12 @@ async def _generate_one(
block_group_request_id = origin_request_id
p_node = None
d_node = None
prefill_load_released = False

try:
p_node, d_node = await self.select_p_d_node(prompt, origin_sampling_params, multimodal_params)
# 记录当前 P 节点的在途 prompt 负载,首 token 生成后即可释放。
p_node.dispatched_prompt_chars += len(prompt)

history_gen_token_strs = []

Expand Down Expand Up @@ -252,6 +255,9 @@ async def _generate_one(
if iter_index == 0 and origin_prompt_cache_len is None:
origin_prompt_cache_len = metadata.get("prompt_cache_len", 0)
metadata["prompt_cache_len"] = origin_prompt_cache_len or 0
if not prefill_load_released:
p_node.dispatched_prompt_chars = max(0, p_node.dispatched_prompt_chars - len(prompt))
prefill_load_released = True
yield origin_request_id, request_output, metadata, finish_status

await self.remove_req(group_request_id=block_group_request_id)
Expand All @@ -271,6 +277,8 @@ async def _generate_one(
raise e

finally:
if p_node is not None and not prefill_load_released:
p_node.dispatched_prompt_chars = max(0, p_node.dispatched_prompt_chars - len(prompt))
await self.remove_req(block_group_request_id)
return

Expand Down Expand Up @@ -731,12 +739,6 @@ def register_pd(self, pd_info_json, websocket):
if pd_client.mode == "prefill":
self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port]
self.prefill_nodes.append(pd_client)
# dispatched_prompt_chars is the cumulative counter used by the
# CacheAware policy to balance request dispatch across prefill nodes.
# Reset all counters together when a node registers or reconnects so
# stale history does not bias CacheAware toward the zero-valued node.
for prefill_node in self.prefill_nodes:
prefill_node.dispatched_prompt_chars = 0
elif pd_client.mode == "decode":
self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != pd_client.client_ip_port]
self.decode_nodes.append(pd_client)
Expand Down
90 changes: 37 additions & 53 deletions lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
- 用前缀树(见 PromptCacheTree)记录「历史 prompt -> 处理它的 worker」;
- 树中的 prefill_node 对应 worker.client_ip_port;
- prompt 会按 sample_stride 抽稀后再插入/匹配,降低树的深度与内存;
- 用 worker.dispatched_prompt_chars(累计派发的 prompt 字符数)做粗粒度均衡
- 用 worker.dispatched_prompt_chars(当前尚未产出首 token 的 prompt 字符数)做负载均衡

选点流程见 CacheAwarePolicy.select_worker。
"""
Expand All @@ -35,8 +35,8 @@ class CacheAwareConfig:

# 前缀匹配成功率阈值:matched_char_count / input_char_count 超过该值才路由到命中节点。
cache_threshold: float = 0.5
# 派发量不均衡判定:max > min * balance_rel_threshold 时强制选派发量最少的节点
balance_rel_threshold: float = 1.2
# cache 命中节点的在途量超过最空闲节点该倍数时,优先选择最空闲节点
balance_rel_threshold: float = 1.8
# 前缀树允许的最大节点数(不含 root)。
max_node_count: int = 1_000_000
# 每次 LRU 驱逐的叶节点数量。
Expand Down Expand Up @@ -66,16 +66,6 @@ def __init__(self, config: Optional[CacheAwareConfig] = None) -> None:
recursion_limit=self.config.recursion_limit,
)

def _select_worker_min_dispatched(
self,
workers: List[PD_Client_Obj],
request_text: str,
) -> PD_Client_Obj:
"""派发量优先兜底:选择累计 dispatched_prompt_chars 最小的 worker,并写入前缀树。"""
min_dispatched_worker = min(workers, key=lambda worker: worker.dispatched_prompt_chars)
self.prompt_cache_tree.insert(request_text, min_dispatched_worker.client_ip_port)
return min_dispatched_worker

def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Optional[PD_Client_Obj]:
"""
为一次请求选择 prefill worker。
Expand All @@ -89,38 +79,19 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti

决策顺序:
1) workers 为空 -> 返回 None;
2) 若 max(dispatched) > min(dispatched) * balance_rel_threshold,
认为派发不均衡,直接选派发量最少的节点;
3) 否则对 request_text 做前缀匹配,计算
2) 对 request_text 做前缀匹配,计算
match_rate = matched_char_count / input_char_count;
4) match_rate > cache_threshold 且命中 prefill_node 仍在线 -> 路由到该节点并更新树;
5) 未命中阈值或 prefill_node 不在当前 workers 中 -> 回退到派发量最少选择。
3) match_rate > cache_threshold 且命中 prefill_node 仍在线 -> 得到 cache 命中节点;
4) cache 命中节点负载未严重高于最空闲节点 -> 选择 cache 命中节点;
5) 未命中或负载严重失衡 -> 选择最空闲节点;
6) 将当前 prompt 与最终选中的节点写入前缀树。
"""
if not workers:
return None
if len(request_text) <= 1:
raise ValueError(f"request_text length must be > 1, got {len(request_text)}")

# ---- 1. 派发均衡门闩:差距过大时不再追求 cache 亲和 ----
dispatched_chars = [worker.dispatched_prompt_chars for worker in workers]
min_dispatched = min(dispatched_chars) if dispatched_chars else 0
max_dispatched = max(dispatched_chars) if dispatched_chars else 0

is_imbalanced = max_dispatched > (min_dispatched * self.config.balance_rel_threshold)

logger.info(
f"CacheAwarePolicy: min_dispatched={min_dispatched}, max_dispatched={max_dispatched}, "
f"balance_rel_threshold={self.config.balance_rel_threshold:.4f}, "
f"is_imbalanced={is_imbalanced}"
)

if is_imbalanced:
return self._select_worker_min_dispatched(
workers=workers,
request_text=request_text,
)

# ---- 2. 前缀匹配:估计当前请求与历史请求的 cache 复用潜力 ----
# ---- 1. 前缀匹配:估计当前请求与历史请求的 cache 复用潜力 ----
result = self.prompt_cache_tree.prefix_match(request_text)
match_rate = 0.0 if result.input_char_count == 0 else result.matched_char_count / result.input_char_count

Expand All @@ -131,26 +102,39 @@ def select_worker(self, workers: List[PD_Client_Obj], request_text: str) -> Opti
f"prefill_node={result.prefill_node}"
)

selected_worker: Optional[PD_Client_Obj] = None
cache_worker: Optional[PD_Client_Obj] = None
if match_rate > self.config.cache_threshold and result.prefill_node is not None:
# 树中的 prefill_node 是 client_ip_port,需要映射回当前在线 worker 对象。
for worker in workers:
if worker.client_ip_port == result.prefill_node:
selected_worker = worker
cache_worker = worker
break

# ---- 2. 负载配平:cache 节点过载时选择当前在途量最少的节点 ----
least_loaded_worker = min(workers, key=lambda worker: worker.dispatched_prompt_chars)
request_load = len(request_text)
least_projected_load = least_loaded_worker.dispatched_prompt_chars + request_load
cache_projected_load = None
cache_worker_is_overloaded = False
if cache_worker is not None:
cache_projected_load = cache_worker.dispatched_prompt_chars + request_load
cache_worker_is_overloaded = cache_projected_load > least_projected_load * self.config.balance_rel_threshold

if cache_worker is None or cache_worker_is_overloaded:
selected_worker = least_loaded_worker
else:
selected_worker = cache_worker

logger.info(
f"CacheAwarePolicy: selected_worker="
f"{selected_worker.client_ip_port if selected_worker else None}, "
f"match_rate={match_rate:.4f}, cache_threshold={self.config.cache_threshold:.4f}"
f"CacheAwarePolicy: cache_worker={cache_worker.client_ip_port if cache_worker else None}, "
f"cache_worker_load={cache_worker.dispatched_prompt_chars if cache_worker else None}, "
f"cache_projected_load={cache_projected_load}, "
f"least_loaded_worker={least_loaded_worker.client_ip_port}, "
f"least_loaded_worker_load={least_loaded_worker.dispatched_prompt_chars}, "
f"least_projected_load={least_projected_load}, "
f"balance_rel_threshold={self.config.balance_rel_threshold:.4f}, "
f"cache_worker_is_overloaded={cache_worker_is_overloaded}, "
f"selected_worker={selected_worker.client_ip_port}"
)

# ---- 3. 命中则更新树;未命中则派发量兜底并写入树 ----
if selected_worker is not None:
self.prompt_cache_tree.insert(request_text, selected_worker.client_ip_port)
return selected_worker
else:
return self._select_worker_min_dispatched(
workers=workers,
request_text=request_text,
)
self.prompt_cache_tree.insert(request_text, selected_worker.client_ip_port)
return selected_worker
Original file line number Diff line number Diff line change
Expand Up @@ -87,9 +87,6 @@ def select_p_d_node(
p_node = self.policy.select_worker(self.prefill_nodes, request_text=prompt)
d_node = self._importance_sampling(self.decode_nodes)

# 累计派发字符数,供后续 cache-aware 做派发均衡判断。
p_node.dispatched_prompt_chars += len(prompt)

logger.info(
f"LoadBalancedCacheAwareSelector: selected p_node={p_node.client_ip_port}, "
f"d_node={d_node.client_ip_port}"
Expand Down
2 changes: 1 addition & 1 deletion lightllm/server/pd_io_struct.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@ class PD_Client_Obj:
start_args: object # 节点的启动参数信息,用于做匹配性的校验,防止运行过程中出现问题。
websocket: WebSocket = None # 用于通信的 websocket 连接对象
run_status: _PD_Client_RunStatus = field(default_factory=_PD_Client_RunStatus)
# cache-aware 选点用:累计派发到该节点的 prompt 字符数(只增不减,非实时负载)
# cache-aware 选点用:当前派发到该节点且尚未产出首 token 的 prompt 字符数。
dispatched_prompt_chars: int = 0

def __post_init__(self):
Expand Down
72 changes: 71 additions & 1 deletion test/test_pd_selector/test_pd_master_multi_choice.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ async def run():
started = set()
all_started = asyncio.Event()
captured_params = []
p_node = MagicMock()
p_node = MagicMock(dispatched_prompt_chars=0)
d_node = MagicMock()
manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node))
manager._split_max_new_tokens = MagicMock(return_value=[4])
Expand Down Expand Up @@ -93,6 +93,7 @@ async def wait_to_token_package(
call("lightllm_request_count"),
call("lightllm_request_success"),
]
assert p_node.dispatched_prompt_chars == 0

asyncio.run(asyncio.wait_for(run(), timeout=2))

Expand Down Expand Up @@ -203,3 +204,72 @@ async def choice():
assert choice_closed.is_set()

asyncio.run(asyncio.wait_for(run(), timeout=2))


def test_pd_master_releases_prefill_load_when_generation_fails():
async def run():
manager = _manager()
manager._split_max_new_tokens = MagicMock(return_value=[4])
manager.id_gen.generate_id.return_value = 808
manager.remove_req = AsyncMock()
manager.abort = AsyncMock()
p_node = MagicMock(dispatched_prompt_chars=0)
d_node = MagicMock()
manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node))

async def failing_wait_to_token_package(*_args, **_kwargs):
raise RuntimeError("generation failed")
yield None

manager._wait_to_token_package = failing_wait_to_token_package

with pytest.raises(RuntimeError, match="generation failed"):
async for _ in manager._generate_one(
"prompt",
SamplingParams(),
MagicMock(),
MagicMock(),
0,
800,
):
pass

assert p_node.dispatched_prompt_chars == 0

asyncio.run(asyncio.wait_for(run(), timeout=2))


def test_pd_master_releases_prefill_load_when_stream_is_closed():
async def run():
manager = _manager()
manager._split_max_new_tokens = MagicMock(return_value=[4])
manager.id_gen.generate_id.return_value = 808
manager.remove_req = AsyncMock()
manager.abort = AsyncMock()
other_request_load = 17
p_node = MagicMock(dispatched_prompt_chars=other_request_load)
d_node = MagicMock()
manager.select_p_d_node = AsyncMock(return_value=(p_node, d_node))

async def wait_to_token_package(*_args, **_kwargs):
yield 808, "first", {"prompt_tokens": 1}, FinishStatus()
await asyncio.sleep(10)

manager._wait_to_token_package = wait_to_token_package
generator = manager._generate_one(
"prompt",
SamplingParams(),
MagicMock(),
MagicMock(),
0,
800,
)

assert (await generator.__anext__())[1] == "first"
assert p_node.dispatched_prompt_chars == other_request_load
await generator.aclose()

assert p_node.dispatched_prompt_chars == other_request_load
manager.abort.assert_awaited_once()

asyncio.run(asyncio.wait_for(run(), timeout=2))
5 changes: 4 additions & 1 deletion unit_tests/server/httpserver/test_pd_master_cached_tokens.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,8 @@ def gen_id():
mgr.metric_client = SimpleNamespace(counter_inc=lambda *a, **k: None, histogram_observe=lambda *a, **k: None)
mgr.tokens = lambda *a, **k: 10
mgr._log_req_header = lambda *a, **k: asyncio.sleep(0)
mgr.select_p_d_node = lambda *a, **k: asyncio.sleep(0, result=(1, 1))
p_node = SimpleNamespace(dispatched_prompt_chars=0)
mgr.select_p_d_node = lambda *a, **k: asyncio.sleep(0, result=(p_node, 1))
mgr.remove_req = lambda *a, **k: asyncio.sleep(0)
return mgr

Expand Down Expand Up @@ -61,6 +62,7 @@ async def run():
def test_single_block_prefill_hit_persists_past_decode_zeros(monkeypatch):
mgr = _make_manager(monkeypatch)
sp = SamplingParams()
sp.n = 1
sp.max_new_tokens = 3
sp.best_of = 1
sp.group_request_id = 0
Expand All @@ -71,6 +73,7 @@ def test_single_block_prefill_hit_persists_past_decode_zeros(monkeypatch):
def test_multi_block_keeps_first_block_hit(monkeypatch):
mgr = _make_manager(monkeypatch)
sp = SamplingParams()
sp.n = 1
sp.max_new_tokens = 5
sp.best_of = 1
sp.group_request_id = 0
Expand Down
46 changes: 46 additions & 0 deletions unit_tests/server/test_pd_cache_aware.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
from types import SimpleNamespace

from lightllm.server.httpserver_for_pd_master.pd_selector.cache_aware import CacheAwarePolicy


def _worker(address: str, dispatched_prompt_chars: int = 0):
return SimpleNamespace(
client_ip_port=address,
dispatched_prompt_chars=dispatched_prompt_chars,
)


def test_cache_aware_keeps_cache_worker_when_inflight_load_is_balanced():
policy = CacheAwarePolicy()
cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=110)
least_loaded_worker = _worker("10.0.0.2:8000", dispatched_prompt_chars=100)
prompt = "shared prefix " * 100
policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port)

selected_worker = policy.select_worker([cache_worker, least_loaded_worker], prompt)

assert selected_worker is cache_worker


def test_cache_aware_uses_least_loaded_worker_when_cache_worker_is_overloaded():
policy = CacheAwarePolicy()
cache_worker = _worker("10.0.0.1:8000", dispatched_prompt_chars=2000)
least_loaded_worker = _worker("10.0.0.2:8000", dispatched_prompt_chars=100)
prompt = "shared prefix " * 100
policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port)

selected_worker = policy.select_worker([cache_worker, least_loaded_worker], prompt)

assert selected_worker is least_loaded_worker


def test_cache_aware_keeps_cache_worker_when_both_workers_are_idle():
policy = CacheAwarePolicy()
cache_worker = _worker("10.0.0.1:8000")
other_worker = _worker("10.0.0.2:8000")
prompt = "shared prefix " * 100
policy.prompt_cache_tree.insert(prompt, cache_worker.client_ip_port)

selected_worker = policy.select_worker([other_worker, cache_worker], prompt)

assert selected_worker is cache_worker
8 changes: 4 additions & 4 deletions unit_tests/server/test_pd_master_mode.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,7 +184,7 @@ def test_pd_manager_without_connected_nodes_is_healthy():
assert asyncio.run(manager.check_pd_nodes_health()) is True


def test_prefill_registration_resets_all_dispatched_prompt_chars():
def test_prefill_registration_preserves_existing_inflight_prompt_chars():
args = StartArgs()
manager = PDManager(args)

Expand All @@ -208,10 +208,10 @@ def register_prefill(node_id, client_ip_port):

register_prefill(2, "10.0.0.2:8000")

assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [0, 0]
assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [1234, 0]


def test_prefill_reconnection_resets_all_dispatched_prompt_chars():
def test_prefill_reconnection_preserves_other_nodes_inflight_prompt_chars():
args = StartArgs()
manager = PDManager(args)

Expand All @@ -235,7 +235,7 @@ def pd_info(node_id, client_ip_port):
manager.register_pd(pd_info(3, "10.0.0.2:8000"), websocket=object())

assert [node.client_ip_port for node in manager.prefill_nodes] == ["10.0.0.1:8000", "10.0.0.2:8000"]
assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [0, 0]
assert [node.dispatched_prompt_chars for node in manager.prefill_nodes] == [100, 0]


def test_pd_master_inference_health_matches_normal_node_semantics(monkeypatch):
Expand Down
Loading