diff --git a/lightllm/server/router/manager.py b/lightllm/server/router/manager.py index b5b29e2383..43d03179c3 100644 --- a/lightllm/server/router/manager.py +++ b/lightllm/server/router/manager.py @@ -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, diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 9a009e0e66..650cb8493c 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -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) diff --git a/unit_tests/common/test_req_manager_page.py b/unit_tests/common/test_req_manager_page.py index cd64346bc1..f45ac26ceb 100644 --- a/unit_tests/common/test_req_manager_page.py +++ b/unit_tests/common/test_req_manager_page.py @@ -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 @@ -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) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py index 02d9d202ab..73e133cc18 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py @@ -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