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
6 changes: 3 additions & 3 deletions lightllm/server/router/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -155,9 +155,9 @@ async def wait_to_model_ready(self):
# The overlapped iteration then needs mtp_step positions for target
# verification and another mtp_step for the DSpark/DFlash draft block.
# Thus the page table needs 3 * mtp_step positions of MTP headroom.
# Keep eight additional positions as a safety margin for future overlap
# changes while preserving the historical +8 for non-MTP runs.
"max_seq_length": self.args.max_req_total_len + 3 * self.args.mtp_step + 8,
# Keep 16 additional positions: eight preserve the historical safety
# margin and eight cover the small-page decode preallocation window.
"max_seq_length": self.args.max_req_total_len + 3 * self.args.mtp_step + 16,
"nccl_host": self.args.nccl_host,
"nccl_port": get_shm_port_args().nccl_port,
"is_first_token_constraint_mode": self.args.first_token_constraint_mode,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -756,6 +756,10 @@ def _get_classed_reqs(
if is_decode:
# KV 容量检查使用额外分配量,已有页的剩余容量可以覆盖部分或全部 decode 需求。
_, alloc_token_num = req_obj.decode_need_token_num()
# page_size 较小时,decode 会频繁触发 KV 内存分配。此处额外预申请不超过 8 个 token,
# 并将数量向下对齐到 page_size 的整数倍,以减少 alloc 调用次数并保持分页分配约束。
if alloc_token_num > 0 and self.args.page_size < 8:
alloc_token_num += 8 // self.args.page_size * self.args.page_size
if alloc_token_num <= can_alloc_token_num:
self._alloc_req_kv_mem(req_obj, alloc_token_num, no_blcoking_copy=True)
decode_reqs.append(req_obj)
Expand Down
55 changes: 52 additions & 3 deletions unit_tests/common/test_req_manager_page.py
Original file line number Diff line number Diff line change
Expand Up @@ -248,7 +248,7 @@ def test_decode_scheduler_uses_non_blocking_request_table_copy(monkeypatch):
backend._timer_merge_radix_tree = lambda: None
backend._reorder_pd_high_priority_reqs = lambda reqs: reqs
backend._reorder_long_prefill_reqs = lambda reqs: reqs
context.get_can_alloc_token_num = lambda: 4
context.get_can_alloc_token_num = lambda: 12
context.cache_placement_controller = SimpleNamespace(set_req_cache_way=lambda reqs: None)
context.filter_reqs = lambda finished_reqs: None
context.pause_reqs = lambda reqs, is_master_in_dp: None
Expand Down Expand Up @@ -286,11 +286,60 @@ def _record_alloc(req_obj, alloc_token_num, no_blcoking_copy=False):

assert prefill_reqs == []
assert decode_reqs == [req]
assert req.hold_kv_len == 8
assert context.req_manager.mem_manager.alloc_sizes == [4]
assert req.hold_kv_len == 16
assert context.req_manager.mem_manager.alloc_sizes == [12]
assert copy_modes == [True]


@pytest.mark.parametrize(
"page_size, expected_alloc_token_num",
[(1, 9), (3, 9), (4, 12), (7, 14), (8, 8)],
)
def test_decode_small_page_preallocation_reuses_one_allocation_for_multiple_steps(
monkeypatch, page_size, expected_alloc_token_num
):
context, backend = _make_context(monkeypatch)
context.args.page_size = page_size
context.req_manager.req_to_token_indexs = torch.full((2, 64), -1, dtype=torch.int32)
backend.args.enable_cpu_cache = False
backend.args.enable_prefill_decode_mixed = False
backend.args.run_mode = "normal"
backend.support_overlap = False
backend.batch_max_tokens = 8
backend.is_master_in_dp = True
backend._timer_merge_radix_tree = lambda: None
backend._reorder_pd_high_priority_reqs = lambda reqs: reqs
backend._reorder_long_prefill_reqs = lambda reqs: reqs
context.get_can_alloc_token_num = lambda: 32
context.cache_placement_controller = SimpleNamespace(set_req_cache_way=lambda reqs: None)
context.filter_reqs = lambda finished_reqs: None
context.pause_reqs = lambda reqs, is_master_in_dp: None

req = _make_req(0)
req.args = context.args
req.cur_kv_len = page_size
req.hold_kv_len = page_size
req.mtp_step = 0
req.filter_mark = False
req.wait_pause = False
req.paused = False
req.infer_aborted = False
req.finish_status = infer_batch.FinishStatus()
req.get_cur_total_len = lambda: req.cur_kv_len + 1
req.decode_need_token_num = MethodType(InferReq.decode_need_token_num, req)
context.req_manager.req_to_token_indexs[0, :page_size] = torch.arange(page_size, dtype=torch.int32)
context.req_manager.mem_manager.next_index = page_size
backend._filter_not_ready_reqs = lambda req_ids: [req]

for _ in range(expected_alloc_token_num):
_, decode_reqs = backend._get_classed_reqs(req_ids=[0])
assert decode_reqs == [req]
req.cur_kv_len += 1

assert req.hold_kv_len == page_size + expected_alloc_token_num
assert context.req_manager.mem_manager.alloc_sizes == [expected_alloc_token_num]


def test_decode_reserves_mtp_headroom(monkeypatch):
context, backend = _make_context(monkeypatch)
req = _make_req(0)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ def _classify_without_token_capacity(monkeypatch, req, support_overlap=True):
enable_cpu_cache=False,
enable_prefill_decode_mixed=False,
run_mode="decode",
page_size=4,
)
backend.support_overlap = support_overlap
backend.is_master_in_dp = True
Expand Down
Loading