From b1623cad3f549413fc704b4d24574190bf23da9c Mon Sep 17 00:00:00 2001 From: Junyi Chen Date: Tue, 4 Aug 2026 20:06:42 +0800 Subject: [PATCH 001/103] feat: add unified LightSpec speculative decoding Unify speculative decoding across MTP, Eagle3, DSpark, and DFlash. Add dynamic verification, MTP-diverse kernels, checkpoint compatibility, and focused GPU and unit test coverage. --- .../common/basemodel/attention/base_att.py | 20 +- lightllm/common/basemodel/attention/fa3/fp.py | 54 +- .../common/basemodel/attention/fa3/fp8.py | 13 +- .../common/basemodel/attention/fa3/mla.py | 51 +- .../common/basemodel/attention/triton/fp.py | 60 +- lightllm/common/basemodel/basemodel.py | 212 ++- lightllm/common/basemodel/batch_objs.py | 36 +- lightllm/common/basemodel/cuda_graph.py | 312 +++- lightllm/common/basemodel/infer_struct.py | 25 +- .../common/basemodel/prefill_cuda_graph.py | 8 +- .../decode_att/gqa/mtp_diverse/__init__.py | 14 + .../gqa/mtp_diverse/mtp_diverse_attn.py | 61 + .../gqa/mtp_diverse/stage1_single_token.py | 333 ++++ .../gqa/mtp_diverse/stage2_single_token.py | 202 ++ .../triton_kernel/dynamic_mtp_utils.py | 180 ++ .../basemodel/triton_kernel/fa3_utils.py | 173 +- .../basemodel/triton_kernel/mtp_utils.py | 441 ++++- lightllm/common/req_manager.py | 17 +- lightllm/common/speculative/__init__.py | 21 + lightllm/common/speculative/config.py | 226 +++ ....bfloat16,q_head_dim=128}_NVIDIA_H800.json | 326 ++++ ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 254 +++ ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 254 +++ ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 26 + ...h.float16,q_head_dim=128}_NVIDIA_H200.json | 26 + ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 26 + ...h.float16,q_head_dim=128}_NVIDIA_H200.json | 26 + ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 26 + ...h.float16,q_head_dim=128}_NVIDIA_H200.json | 26 + ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 26 + ...h.float16,q_head_dim=128}_NVIDIA_H200.json | 26 + lightllm/models/__init__.py | 219 ++- lightllm/models/deepseek_mtp/model.py | 7 +- lightllm/models/glm4_moe_lite_mtp/model.py | 7 +- lightllm/models/mistral_mtp/model.py | 7 +- lightllm/models/qwen3_dflash/__init__.py | 3 + lightllm/models/qwen3_dflash/infer_struct.py | 40 + .../qwen3_dflash/layer_infer/__init__.py | 9 + .../layer_infer/post_layer_infer.py | 28 + .../layer_infer/pre_layer_infer.py | 67 + .../layer_infer/transformer_layer_infer.py | 113 ++ .../qwen3_dflash/layer_weights/__init__.py | 11 + .../pre_and_post_layer_weight.py | 63 + .../layer_weights/transformer_layer_weight.py | 102 ++ lightllm/models/qwen3_dflash/model.py | 159 ++ lightllm/models/qwen3_dspark/__init__.py | 3 + .../qwen3_dspark/layer_infer/__init__.py | 4 + .../layer_infer/post_layer_infer.py | 225 +++ .../qwen3_dspark/layer_weights/__init__.py | 5 + .../pre_and_post_layer_weight.py | 67 + lightllm/models/qwen3_dspark/model.py | 15 + lightllm/models/qwen3_eagle/__init__.py | 0 .../qwen3_eagle/layer_infer/__init__.py | 7 + .../layer_infer/pre_layer_infer.py | 55 + .../layer_infer/transformer_layer_infer.py | 67 + .../qwen3_eagle/layer_weights/__init__.py | 0 .../pre_and_post_layer_weight.py | 64 + .../layer_weights/transformer_layer_weight.py | 111 ++ lightllm/models/qwen3_eagle/model.py | 85 + lightllm/models/qwen3_moe_mtp/model.py | 7 +- lightllm/server/api_cli.py | 52 +- lightllm/server/api_start.py | 50 +- lightllm/server/core/objs/req.py | 5 +- lightllm/server/core/objs/start_args_type.py | 11 +- lightllm/server/httpserver/manager.py | 50 +- .../httpserver_for_pd_master/manager.py | 10 +- .../server/router/model_infer/infer_batch.py | 58 +- .../model_infer/mode_backend/__init__.py | 61 +- .../model_infer/mode_backend/base_backend.py | 223 ++- .../mode_backend/chunked_prefill/impl.py | 252 +-- .../mode_backend/diverse_backend/impl.py | 6 +- .../mode_backend/dp_backend/impl.py | 330 ++-- .../mode_backend/generic_post_process.py | 193 +- .../mode_backend/generic_pre_process.py | 39 + .../mode_backend/update_mem_index.py | 50 + .../model_infer/speculative/__init__.py | 42 + .../router/model_infer/speculative/planner.py | 1626 +++++++++++++++++ .../speculative/proposers/__init__.py | 61 + .../model_infer/speculative/proposers/base.py | 191 ++ .../speculative/proposers/dflash.py | 189 ++ .../speculative/proposers/dspark.py | 126 ++ .../speculative/proposers/eagle3.py | 210 +++ .../speculative/proposers/eagle_mtp.py | 219 +++ .../speculative/proposers/vanilla_mtp.py | 119 ++ .../router/model_infer/speculative/runner.py | 213 +++ .../router/model_infer/speculative/runtime.py | 774 ++++++++ .../router/model_infer/speculative/state.py | 130 ++ .../model_infer/speculative/verifier.py | 107 ++ lightllm/utils/envs_utils.py | 43 +- lightllm/utils/kv_cache_utils.py | 7 +- .../gqa/mtp_diverse/test_mtp_diverse.py | 195 ++ .../test_int8kv_flash_decoding_diverse.py | 12 +- .../triton_kernel/test_dynamic_mtp_utils.py | 232 +++ .../basemodel/triton_kernel/test_fa3_utils.py | 82 + .../basemodel/triton_kernel/test_mtp_utils.py | 166 ++ unit_tests/common/speculative/test_config.py | 40 + .../mode_backend/test_generic_post_process.py | 50 + .../model_infer/speculative/test_planner.py | 66 + .../server/test_api_start_spec_config.py | 42 + 99 files changed, 10616 insertions(+), 767 deletions(-) create mode 100644 lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/__init__.py create mode 100644 lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/mtp_diverse_attn.py create mode 100644 lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/stage1_single_token.py create mode 100644 lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/stage2_single_token.py create mode 100644 lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py create mode 100644 lightllm/common/speculative/__init__.py create mode 100644 lightllm/common/speculative/config.py create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H800/_fwd_kernel_mtp_diverse_stage1_single_token:v1/{block_seq=256,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=8,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json create mode 100644 lightllm/models/qwen3_dflash/__init__.py create mode 100644 lightllm/models/qwen3_dflash/infer_struct.py create mode 100644 lightllm/models/qwen3_dflash/layer_infer/__init__.py create mode 100644 lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py create mode 100644 lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py create mode 100644 lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py create mode 100644 lightllm/models/qwen3_dflash/layer_weights/__init__.py create mode 100644 lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py create mode 100644 lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py create mode 100644 lightllm/models/qwen3_dflash/model.py create mode 100644 lightllm/models/qwen3_dspark/__init__.py create mode 100644 lightllm/models/qwen3_dspark/layer_infer/__init__.py create mode 100644 lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py create mode 100644 lightllm/models/qwen3_dspark/layer_weights/__init__.py create mode 100644 lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py create mode 100644 lightllm/models/qwen3_dspark/model.py create mode 100644 lightllm/models/qwen3_eagle/__init__.py create mode 100644 lightllm/models/qwen3_eagle/layer_infer/__init__.py create mode 100644 lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py create mode 100644 lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py create mode 100644 lightllm/models/qwen3_eagle/layer_weights/__init__.py create mode 100644 lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py create mode 100644 lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py create mode 100644 lightllm/models/qwen3_eagle/model.py create mode 100644 lightllm/server/router/model_infer/mode_backend/update_mem_index.py create mode 100644 lightllm/server/router/model_infer/speculative/__init__.py create mode 100644 lightllm/server/router/model_infer/speculative/planner.py create mode 100644 lightllm/server/router/model_infer/speculative/proposers/__init__.py create mode 100644 lightllm/server/router/model_infer/speculative/proposers/base.py create mode 100644 lightllm/server/router/model_infer/speculative/proposers/dflash.py create mode 100644 lightllm/server/router/model_infer/speculative/proposers/dspark.py create mode 100644 lightllm/server/router/model_infer/speculative/proposers/eagle3.py create mode 100644 lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py create mode 100644 lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py create mode 100644 lightllm/server/router/model_infer/speculative/runner.py create mode 100644 lightllm/server/router/model_infer/speculative/runtime.py create mode 100644 lightllm/server/router/model_infer/speculative/state.py create mode 100644 lightllm/server/router/model_infer/speculative/verifier.py create mode 100644 unit_tests/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/test_mtp_diverse.py create mode 100644 unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py create mode 100644 unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py create mode 100644 unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py create mode 100644 unit_tests/common/speculative/test_config.py create mode 100644 unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py create mode 100644 unit_tests/server/router/model_infer/speculative/test_planner.py create mode 100644 unit_tests/server/test_api_start_spec_config.py diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index 55d97d2aa8..a1985a6fb9 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -1,7 +1,7 @@ import torch from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import Optional, TYPE_CHECKING, Tuple, Union, Dict +from typing import TYPE_CHECKING, Tuple, Union, Dict if TYPE_CHECKING: from lightllm.common.basemodel.basemodel import TpPartBaseModel @@ -10,23 +10,23 @@ class BaseAttBackend: """ - 用于创建支持各种不同的AttBackend, 如 fa3, flashinfer, triton 实现等, - 这个是单列模式, 每种backend只有一个实例 + 用于创建支持各种不同的AttBackend, 如 fa3, flashinfer, triton 实现等。 + 每个 model 复用一个 backend 实例。 """ _instances = {} def __new__(cls, *args, **kwargs): """ - 重写__new__方法实现单例模式 + Main 和 speculative draft model 可能使用不同的 CUDA graph 上限 + 和缓存布局,不能只按 backend class 共享实例。 """ - # 检查是否已经有该类的实例 - if cls not in cls._instances: - # 创建新实例并存储 + model = kwargs.get("model", args[0] if args else None) + instance_key = (cls, id(model)) + if instance_key not in cls._instances: instance = super().__new__(cls) - cls._instances[cls] = instance - # 返回已有的实例 - return cls._instances[cls] + cls._instances[instance_key] = instance + return cls._instances[instance_key] def __init__(self, model: "TpPartBaseModel"): self.model = model diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index 57f3ab6fe3..76cdd83044 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -1,11 +1,13 @@ import dataclasses import torch from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl -from typing import Optional, TYPE_CHECKING from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.sgl_utils import flash_attn_with_kvcache, flash_attn_with_kvcache_autotune -from lightllm.utils.envs_utils import get_env_start_args -from lightllm.common.basemodel.triton_kernel.fa3_utils import page_table_copy +from lightllm.utils.envs_utils import enable_dynamic_mtp_verify, get_env_start_args +from lightllm.common.basemodel.triton_kernel.fa3_utils import ( + build_dynamic_mtp_fa3_decode_params, + page_table_copy, +) from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor @@ -102,7 +104,7 @@ def _nomarl_prefill_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.infer_state.max_q_seq_len, softmax_scale=sm_scale, - causal=True, + causal=getattr(self.infer_state, "prefill_causal", True), window_size=window_size, softcap=0.0, k_descale=k_descale, @@ -124,9 +126,32 @@ class Fa3DecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: Fa3AttBackend = self.backend - args_mtp_step = get_env_start_args().mtp_step + args_mtp_step = getattr(self.infer_state, "decode_mtp_step", None) + if args_mtp_step is None: + args_mtp_step = get_env_start_args().mtp_step + if self.infer_state.disable_mtp_decode_att: + args_mtp_step = 0 + is_block_draft_decode = getattr(self.infer_state, "is_draft_model", False) + is_dynamic_mtp = ( + args_mtp_step > 0 + and enable_dynamic_mtp_verify() + and not is_block_draft_decode + and not getattr(self.infer_state, "use_static_mtp_layout", False) + ) - if args_mtp_step > 0: + if is_dynamic_mtp: + att_batch_size = self.infer_state.batch_size + (b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len,) = build_dynamic_mtp_fa3_decode_params( + b_req_idx=self.infer_state.b_req_idx, + b_seq_len=self.infer_state.b_seq_len, + b_mark_shared_group=self.infer_state.b_mark_shared_group, + att_batch_size=att_batch_size, + hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, + ) + b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) + self.cu_seqlens_q = b1_cu_q_seq_len.int() + self.cu_seqlens_k = b1_cu_kv_seq_len.int() + elif args_mtp_step > 0: # 修正 mtp 在 fa3 下的输入。 mtp_size = args_mtp_step + 1 b_q_seq_len = torch.full( @@ -139,12 +164,14 @@ def init_state(self): b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) self.cu_seqlens_q = b1_cu_q_seq_len.int() self.cu_seqlens_k = b1_cu_kv_seq_len.int() + att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) else: self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() + att_batch_size = self.infer_state.batch_size - att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) - assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 + if not is_dynamic_mtp: + assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 model = self.backend.model # 可以使用 cuda graph的时候从 buffer中申请 @@ -163,7 +190,14 @@ def init_state(self): device=self.infer_state.input_ids.device, ) - if args_mtp_step > 0: + if is_dynamic_mtp: + page_table_copy( + page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], + req_to_token_indexs=model.req_manager.req_to_token_indexs, + b_req_idx=b_att_req_idx, + ) + self.decode_max_q_seq_len = args_mtp_step + 1 + elif args_mtp_step > 0: page_table_copy( page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], req_to_token_indexs=model.req_manager.req_to_token_indexs, @@ -232,7 +266,7 @@ def _normal_decode_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, softmax_scale=sm_scale, - causal=True, + causal=getattr(self.infer_state, "decode_causal", True), window_size=window_size, softcap=0.0, k_descale=k_descale, diff --git a/lightllm/common/basemodel/attention/fa3/fp8.py b/lightllm/common/basemodel/attention/fa3/fp8.py index adc8b5c01e..dd0222ff5b 100644 --- a/lightllm/common/basemodel/attention/fa3/fp8.py +++ b/lightllm/common/basemodel/attention/fa3/fp8.py @@ -1,12 +1,9 @@ import dataclasses import torch from ..base_att import AttControl -from typing import Optional, TYPE_CHECKING from lightllm.utils.sgl_utils import flash_attn_with_kvcache -from lightllm.utils.envs_utils import get_env_start_args from lightllm.common.basemodel.triton_kernel.quantization.q_per_head_fp8_quant import q_per_head_fp8_quant from lightllm.utils.vllm_utils import HAS_VLLM, vllm_ops -from typing import Union from .fp import Fa3AttBackend, Fa3PrefillAttState, Fa3DecodeAttState if HAS_VLLM: @@ -99,7 +96,7 @@ def _fp8_prefill_att( cu_seqlens_q=self.cu_seqlens_q, cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.infer_state.max_q_seq_len, - causal=True, + causal=getattr(self.infer_state, "prefill_causal", True), window_size=(-1, -1), softcap=0.0, q_descale=q_scale, @@ -119,11 +116,7 @@ def init_state(self): super().init_state() self.backend: Fp8Fa3AttBackend = self.backend - args_mtp_step = get_env_start_args().mtp_step - att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) - assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 - - batch_size = att_batch_size + batch_size = self.page_table.shape[0] mem_manager = self.backend.model.mem_manager offline_scales: torch.Tensor = mem_manager.scales @@ -190,7 +183,7 @@ def _fp8_decode_att( cu_seqlens_q=self.cu_seqlens_q, cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, - causal=True, + causal=getattr(self.infer_state, "decode_causal", True), window_size=(-1, -1), softcap=0.0, q_descale=q_scale.view(self.infer_state.batch_size, k_head_num), diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index 9a10457b12..03c277c16c 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -1,11 +1,11 @@ import dataclasses import torch from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl -from typing import Optional, TYPE_CHECKING, Tuple +from typing import Tuple from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.sgl_utils import flash_attn_with_kvcache -from lightllm.utils.envs_utils import get_env_start_args -from lightllm.common.basemodel.triton_kernel.fa3_utils import page_table_copy +from lightllm.utils.envs_utils import enable_dynamic_mtp_verify, get_env_start_args +from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_mtp_fa3_decode_params, page_table_copy from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor from lightllm.utils.sgl_utils import flash_attn_varlen_func @@ -90,7 +90,7 @@ def _mla_prefill_att( max_seqlen_q=self.infer_state.max_q_seq_len, max_seqlen_k=self.infer_state.max_kv_seq_len, softmax_scale=softmax_scale, - causal=True, + causal=getattr(self.infer_state, "prefill_causal", True), return_softmax_lse=False, ) return o_tensor @@ -107,9 +107,32 @@ class MlaFa3DecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: MlaFa3AttBackend = self.backend - args_mtp_step = get_env_start_args().mtp_step + args_mtp_step = getattr(self.infer_state, "decode_mtp_step", None) + if args_mtp_step is None: + args_mtp_step = get_env_start_args().mtp_step + if self.infer_state.disable_mtp_decode_att: + args_mtp_step = 0 + is_block_draft_decode = getattr(self.infer_state, "is_draft_model", False) + is_dynamic_mtp = ( + args_mtp_step > 0 + and enable_dynamic_mtp_verify() + and not is_block_draft_decode + and not getattr(self.infer_state, "use_static_mtp_layout", False) + ) - if args_mtp_step > 0: + if is_dynamic_mtp: + att_batch_size = self.infer_state.batch_size + (b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len,) = build_dynamic_mtp_fa3_decode_params( + b_req_idx=self.infer_state.b_req_idx, + b_seq_len=self.infer_state.b_seq_len, + b_mark_shared_group=self.infer_state.b_mark_shared_group, + att_batch_size=att_batch_size, + hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, + ) + b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) + self.cu_seqlens_q = b1_cu_q_seq_len.int() + self.cu_seqlens_k = b1_cu_kv_seq_len.int() + elif args_mtp_step > 0: # 修正 mtp 在 fa3 下的输入。 mtp_size = args_mtp_step + 1 b_q_seq_len = torch.full( @@ -126,8 +149,9 @@ def init_state(self): self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() - att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) - assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 + if not is_dynamic_mtp: + assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 + att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) model = self.backend.model # 可以使用 cuda graph的时候从 buffer中申请 @@ -146,7 +170,14 @@ def init_state(self): device=self.infer_state.input_ids.device, ) - if args_mtp_step > 0: + if is_dynamic_mtp: + page_table_copy( + page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], + req_to_token_indexs=model.req_manager.req_to_token_indexs, + b_req_idx=b_att_req_idx, + ) + self.decode_max_q_seq_len = args_mtp_step + 1 + elif args_mtp_step > 0: page_table_copy( page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], req_to_token_indexs=model.req_manager.req_to_token_indexs, @@ -219,7 +250,7 @@ def _mla_decode_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, softmax_scale=softmax_scale, - causal=True, + causal=getattr(self.infer_state, "decode_causal", True), window_size=(-1, -1), softcap=0.0, k_descale=k_descale, diff --git a/lightllm/common/basemodel/attention/triton/fp.py b/lightllm/common/basemodel/attention/triton/fp.py index a1370a7045..9e1407b483 100644 --- a/lightllm/common/basemodel/attention/triton/fp.py +++ b/lightllm/common/basemodel/attention/triton/fp.py @@ -1,7 +1,8 @@ import dataclasses import torch + +from lightllm.utils.envs_utils import enable_dynamic_mtp_verify, get_env_start_args, enable_triton_mtp_kernel from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl -from typing import Optional class TritonAttBackend(BaseAttBackend): @@ -93,8 +94,21 @@ def _nomarl_prefill_att( @dataclasses.dataclass class TritonDecodeAttState(BaseDecodeAttState): + # MTP related state variables + b_mark_shared_group: torch.Tensor = None + def init_state(self): - pass + args_mtp_step = getattr(self.infer_state, "decode_mtp_step", None) + if args_mtp_step is None: + args_mtp_step = get_env_start_args().mtp_step + if self.infer_state.disable_mtp_decode_att: + args_mtp_step = 0 + + if args_mtp_step > 0: + # MTP mode initialization + self.b_mark_shared_group = self.infer_state.b_mark_shared_group + else: + self.b_mark_shared_group = None def copy_for_decode_cuda_graph(self, new_state: "TritonDecodeAttState"): super().copy_for_decode_cuda_graph(new_state) @@ -112,10 +126,20 @@ def decode_att( assert att_control.tp_alibi is not None return self._alibi_decode_att(q=q, k=k, v=v, att_control=att_control, alloc_func=alloc_func) else: + args_mtp_step = getattr(self.infer_state, "decode_mtp_step", None) + if args_mtp_step is None: + args_mtp_step = get_env_start_args().mtp_step + if self.infer_state.disable_mtp_decode_att: + args_mtp_step = 0 + q_head_num = q.shape[1] k_head_num = k.shape[1] - if q_head_num == k_head_num: - assert att_control.use_sliding_window is False, "sliding_window not supported in non-gqa attention yet" + + if args_mtp_step > 0 and (enable_dynamic_mtp_verify() or enable_triton_mtp_kernel()): + # MTP mode: use mtp diverse attention + assert q_head_num >= k_head_num, "MTP diverse attention requires q_head_num >= k_head_num" + return self._dynamic_mtp_decode_gqa_att(q=q, k=k, v=v, alloc_func=alloc_func) + elif q_head_num == k_head_num: return self._normal_decode_flash_decoding_att(q=q, k=k, v=v, alloc_func=alloc_func) elif q_head_num > k_head_num: return self._normal_decode_gqa_flash_decoding_att( @@ -205,6 +229,34 @@ def _normal_decode_gqa_flash_decoding_att( return out + def _dynamic_mtp_decode_gqa_att( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + alloc_func=torch.empty, + ): + from ...triton_kernel.att.decode_att.gqa.mtp_diverse import ( + token_decode_attention_mtp_diverse_single_token, + ) + + b_seq_len = self.infer_state.b_seq_len + # 在动态 MTP 验证模式下,使用 infer_state.b_mark_shared_group(从 model_input 传递) + # 在静态 MTP 模式下,使用 self.b_mark_shared_group(在 init_state 中初始化) + b_mark_shared_group = self.infer_state.b_mark_shared_group + out = token_decode_attention_mtp_diverse_single_token( + q=q, + k=k, + v=v, + Req_to_tokens=self.infer_state.req_manager.req_to_token_indexs, + B_req_idx=self.infer_state.b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + alloc_tensor_func=alloc_func, + ) + + return out + def _normal_decode_gqa_flash_decoding_att_vsm( self, q: torch.Tensor, diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 0f1bfa9cc6..ab97b81469 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -14,7 +14,6 @@ from lightllm.common.basemodel.infer_struct import InferStateInfo from lightllm.common.kv_cache_mem_manager import MemoryManager from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class -from lightllm.common.req_manager import ReqManager from lightllm.common.infer_utils import init_req_to_token_indexes from lightllm.common.build_utils import repair_config from lightllm.common.basemodel.triton_kernel.copy_kv_index_to_req import copy_kv_index_to_req @@ -25,7 +24,12 @@ from lightllm.common.basemodel.triton_kernel.gather_token_id import gather_token, gather_token_prefill_decode_mixed from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_dp_world_size -from lightllm.utils.envs_utils import get_env_start_args, get_llm_data_type, get_added_mtp_kv_layer_num +from lightllm.utils.envs_utils import ( + enable_triton_mtp_kernel, + get_env_start_args, + get_llm_data_type, + get_added_mtp_kv_layer_num, +) from lightllm.distributed.communication_op import dist_group_manager from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.utils.custom_kernel_utis import pad2dim_tensor_to_new_batch @@ -33,6 +37,7 @@ set_model_init_status, enable_diverse_mode_gqa_decode_fast_kernel, enable_full_att_decode_tune, + enable_dynamic_mtp_verify, ) from lightllm.common.triton_utils.autotuner import Autotuner from lightllm.utils.infer_utils import post_empty_cache @@ -87,6 +92,10 @@ def __init__(self, kvargs): ) # mtp 模式下需要修缮对应的最大batch size,为 (mtp_step + 1) 的倍数 self.graph_max_batch_size = self.graph_max_batch_size * (mtp_step + 1) + self.graph_split_batch_size = int(kvargs.get("graph_split_batch_size", self.args.graph_split_batch_size)) + self.graph_grow_step_size = int(kvargs.get("graph_grow_step_size", self.args.graph_grow_step_size)) + assert self.graph_split_batch_size > 0 + assert self.graph_grow_step_size > 0 self.graph_max_len_in_batch = kvargs.get("graph_max_len_in_batch", 8192) self.disable_cudagraph = kvargs.get("disable_cudagraph", False) @@ -98,15 +107,11 @@ def __init__(self, kvargs): self.enable_tpsp_mix_mode = get_env_start_args().enable_tpsp_mix_mode self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) - self.is_mtp_mode = self.args.mtp_mode in [ - "vanilla_with_att", - "eagle_with_att", - "vanilla_no_att", - "eagle_no_att", - ] self.prefill_graph: PrefillCudaGraph = None + self.spec_adapter = None self._init_config() + self._verify_must() self._verify_params() self._init_quant() @@ -146,6 +151,17 @@ def __init__(self, kvargs): set_model_init_status(True) return + def _wait_other_modules_ready(self): + for event in self.wait_events: + event.wait() + return + + def set_spec_adapter(self, spec_adapter): + self.spec_adapter = spec_adapter + if self.graph is not None: + self.graph.set_spec_adapter(spec_adapter, model=self) + return + def _init_config(self): with open(os.path.join(self.weight_dir_, "config.json"), "r") as json_file: self.config = json.load(json_file) @@ -155,6 +171,12 @@ def _init_config(self): repair_config(self.config, same_names=["num_hidden_layers", "n_layer"]) if self.finetune_config: self.config["vocab_size"] = self.finetune_config.vocab_size + + # eagle3 mode 下,需要修改 vocab_size 为 draft_vocab_size, 其他场景 + # 这个代码并不会生效。 + if "draft_vocab_size" in self.config.keys(): + self.config["target_vocab_size"] = self.config["vocab_size"] + self.config["vocab_size"] = self.config["draft_vocab_size"] return @final @@ -231,6 +253,8 @@ def _check_mem_size(self): return def _init_req_manager(self): + from lightllm.common.req_manager import ReqManager + create_max_seq_len = 0 if self.batch_max_tokens is not None: @@ -279,6 +303,8 @@ def _init_cudagraph(self): max_batch_size=self.graph_max_batch_size, max_len_in_batch=self.graph_max_len_in_batch, tp_world_size=self.tp_world_size_, + graph_split_batch_size=self.graph_split_batch_size, + graph_grow_step_size=self.graph_grow_step_size, ) ) if self.graph is not None: @@ -384,6 +410,8 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) if enable_diverse_mode_gqa_decode_fast_kernel(): infer_state.b_shared_seq_len = model_input.b_shared_seq_len infer_state.b_mark_shared_group = model_input.b_mark_shared_group + elif enable_dynamic_mtp_verify() or enable_triton_mtp_kernel(): + infer_state.b_mark_shared_group = model_input.b_mark_shared_group infer_state.multimodal_params = model_input.multimodal_params @@ -396,6 +424,8 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) # 特殊模型,特殊模式的特定变量初始化操作。 infer_state.mtp_draft_input_hiddens = model_input.mtp_draft_input_hiddens + infer_state.disable_mtp_decode_att = model_input.disable_mtp_decode_att + infer_state.use_static_mtp_layout = model_input.use_static_mtp_layout if infer_state.is_prefill: infer_state.prefill_att_state = self.prefill_att_backend.create_att_prefill_state(infer_state=infer_state) @@ -454,6 +484,11 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.b_mark_shared_group = F.pad( new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=1 ) + elif enable_dynamic_mtp_verify() or enable_triton_mtp_kernel(): + assert new_model_input.b_mark_shared_group is not None + new_model_input.b_mark_shared_group = F.pad( + new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=0 + ) # 特殊模型,特殊模式的特殊变量的特殊 padding if new_model_input.mtp_draft_input_hiddens is not None: @@ -512,22 +547,46 @@ def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle new_model_input.check_input() return new_model_input - def _create_unpad_decode_model_output(self, model_output: ModelOutput, origin_batch_size: int): + def _create_unpad_decode_model_output( + self, + model_output: ModelOutput, + origin_batch_size: int, + microbatch_index: int = 0, + ): padded_batch_size = model_output.logits.shape[0] if padded_batch_size == origin_batch_size: return model_output new_model_output = copy.copy(model_output) new_model_output.logits = new_model_output.logits[0:origin_batch_size] - - # 特殊模型,特殊模式的特殊变量的特殊 unpad - if new_model_output.mtp_main_output_hiddens is not None: - _hidden_states = new_model_output.mtp_main_output_hiddens - new_model_output.mtp_main_output_hiddens = _hidden_states[0:origin_batch_size] + if new_model_output.mtp_draft_confidence_logits is not None: + confidence_rows = new_model_output.mtp_draft_confidence_logits.shape[0] + if confidence_rows == padded_batch_size: + confidence_origin_rows = origin_batch_size + else: + assert padded_batch_size % confidence_rows == 0, ( + "padded decode logits rows must be divisible by confidence rows: " + f"{padded_batch_size}, got {confidence_rows}" + ) + rows_per_confidence = padded_batch_size // confidence_rows + assert origin_batch_size % rows_per_confidence == 0, ( + "origin decode rows must align with confidence row grouping: " + f"{origin_batch_size}, rows_per_confidence={rows_per_confidence}" + ) + confidence_origin_rows = origin_batch_size // rows_per_confidence + new_model_output.mtp_draft_confidence_logits = new_model_output.mtp_draft_confidence_logits[ + 0:confidence_origin_rows + ] + if self.spec_adapter is not None: + self.spec_adapter.unpad_hidden(token_num=origin_batch_size, microbatch_index=microbatch_index) return new_model_output def _create_unpad_prefill_model_output( - self, padded_model_output: ModelOutput, origin_handle_token_num: int, origin_batch_size: int + self, + padded_model_output: ModelOutput, + origin_handle_token_num: int, + origin_batch_size: int, + microbatch_index: int = 0, ): new_model_output = copy.copy(padded_model_output) # logits 始终只对应每个请求最后一个位置,移除 padding 的 req 对应的行。 @@ -541,6 +600,8 @@ def _create_unpad_prefill_model_output( if new_model_output.mtp_main_output_hiddens is not None: _hidden_states = new_model_output.mtp_main_output_hiddens new_model_output.mtp_main_output_hiddens = _hidden_states[0:origin_handle_token_num] + if self.spec_adapter is not None: + self.spec_adapter.unpad_hidden(token_num=origin_handle_token_num, microbatch_index=microbatch_index) return new_model_output @@ -619,14 +680,15 @@ def _decode( infer_batch_size = model_input.batch_size if self.graph is not None and self.graph.can_run( - batch_size=infer_batch_size, max_len_in_batch=model_input.max_kv_seq_len + batch_size=infer_batch_size, + max_len_in_batch=model_input.max_kv_seq_len, ): infer_batch_size = self.graph.find_closest_graph_batch_size(batch_size=infer_batch_size) model_input = self._create_padded_decode_model_input( model_input=model_input, new_batch_size=infer_batch_size ) + need_capture = self.graph.need_capture(infer_batch_size, model_context=model_input) infer_state = self._create_inferstate(model_input) - need_capture = self.graph.need_capture(infer_batch_size) infer_state.is_cuda_graph = need_capture copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -674,10 +736,23 @@ def _context_forward(self, infer_state: InferStateInfo): input_tensors = [input_embs] def prefill_func(input_tensors, infer_state): + spec_context = ( + self.spec_adapter.create_forward_context(self, infer_state) if self.spec_adapter is not None else None + ) _input_embs = input_tensors[0] for i in range(self.layers_num): layer = self.layers_infer[i] _input_embs = layer.context_forward(_input_embs, infer_state, self.trans_layers_weight[i]) + if spec_context is not None: + spec_context.add_hidden( + layer_index=i, + layer_num=self.layers_num, + hidden=_input_embs, + ) + + layer_hidden = spec_context.build_layer_hidden() if spec_context is not None else None + if layer_hidden is not None: + return [_input_embs, layer_hidden] return [_input_embs] handle_token_num = infer_state.input_ids.shape[0] @@ -712,12 +787,18 @@ def prefill_func(input_tensors, infer_state): predict_logits = self.post_infer.token_forward(last_input_embs, infer_state, self.pre_post_weight) model_output = ModelOutput(logits=predict_logits, prompt_logics=infer_state.prompt_logics) - # 特殊模型特殊模式的额外输出 - if self.is_mtp_mode: - input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) - if infer_state.need_dp_prefill_balance: - input_embs = infer_state._all_to_all_unbalance_get(data=input_embs) - model_output.mtp_main_output_hiddens = input_embs.contiguous() + if self.spec_adapter is not None: + if len(output_tensors) > 1: + spec_hidden = self.pre_infer._tpsp_allgather(input=output_tensors[1], infer_state=infer_state) + if infer_state.need_dp_prefill_balance: + spec_hidden = infer_state._all_to_all_unbalance_get(data=spec_hidden) + else: + spec_hidden = last_input_embs + self.spec_adapter.capture_hidden( + infer_state=infer_state, + hidden=spec_hidden.contiguous(), + final_hidden=last_input_embs.contiguous(), + ) # 在开启使用deepep的时候,需要调用clear_deepep_buffer做资源清理,没有启用的时候 # 该调用没有实际意义 @@ -731,21 +812,40 @@ def _token_forward(self, infer_state: InferStateInfo): input_embs = self.pre_infer.token_forward(cuda_input_ids, infer_state, self.pre_post_weight) input_embs = self.pre_infer._tpsp_sp_split(input=input_embs, infer_state=infer_state) + spec_context = ( + self.spec_adapter.create_forward_context(self, infer_state) if self.spec_adapter is not None else None + ) for i in range(self.layers_num): layer = self.layers_infer[i] input_embs: torch.Tensor = layer.token_forward(input_embs, infer_state, self.trans_layers_weight[i]) + if spec_context is not None: + spec_context.add_hidden(layer_index=i, layer_num=self.layers_num, hidden=input_embs) last_input_embs = self.post_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) predict_logits: torch.Tensor = self.post_infer.token_forward( last_input_embs, infer_state=infer_state, layer_weight=self.pre_post_weight ) - model_output = ModelOutput(logits=predict_logits.contiguous()) + mtp_draft_confidence_logits = None + pop_confidence_logits = getattr(self.post_infer, "pop_mtp_draft_confidence_logits", None) + if pop_confidence_logits is not None: + mtp_draft_confidence_logits = pop_confidence_logits() - # 特殊模型特殊模式的额外输出 - if self.is_mtp_mode: - input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) - model_output.mtp_main_output_hiddens = input_embs.contiguous() + model_output = ModelOutput( + logits=predict_logits.contiguous(), + mtp_draft_confidence_logits=mtp_draft_confidence_logits, + ) + + if spec_context is not None: + spec_hidden = spec_context.build_layer_hidden() + if spec_hidden is not None: + spec_hidden = self.pre_infer._tpsp_allgather(input=spec_hidden, infer_state=infer_state) + else: + spec_hidden = last_input_embs + spec_context.capture( + hidden=spec_hidden.contiguous(), + final_hidden=last_input_embs.contiguous(), + ) # 在 cuda graph 模式下,输出需要转为 no ref tensor, 加强mem pool 的复用,降低显存的使用。 if infer_state.is_cuda_graph: @@ -831,11 +931,13 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod padded_model_output=model_output0, origin_handle_token_num=origin_handle_token_num0, origin_batch_size=origin_batch_size0, + microbatch_index=0, ) model_output1 = self._create_unpad_prefill_model_output( padded_model_output=model_output1, origin_handle_token_num=origin_handle_token_num1, origin_batch_size=origin_batch_size1, + microbatch_index=1, ) # 在开启使用deepep的时候,需要调用clear_deepep_buffer做资源清理,没有启用的时候 # 该调用没有实际意义 @@ -873,12 +975,17 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode if self.graph is not None and self.graph.can_run(infer_batch_size, max_len_in_batch): infer_batch_size = self.graph.find_closest_graph_batch_size(infer_batch_size) - need_capture = self.graph.need_capture(infer_batch_size) # TODO 如果支持动态步数的 mtp,在不同的mtp步上,model_input0 和 model_input1 的内部batch size可能不 # 一致,需要按照较高 batch size 进行graph的寻找,同时,进行有效的恢复。 padded_model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) padded_model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) + need_capture = self.graph.need_capture( + infer_batch_size, + model_context=padded_model_input0, + model_context1=padded_model_input1, + ) infer_state0 = self._create_inferstate(padded_model_input0, 0) + infer_state1 = self._create_inferstate(padded_model_input1, 1) infer_state0.is_cuda_graph = need_capture copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -889,7 +996,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state0.init_some_extra_state(self) infer_state0.init_att_state() - infer_state1 = self._create_inferstate(padded_model_input1, 1) infer_state1.is_cuda_graph = need_capture copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -913,8 +1019,12 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode ) # TODO 动态 mtp fix - model_output0 = self._create_unpad_decode_model_output(model_output0, origin_batch_size=origin_batch_size) - model_output1 = self._create_unpad_decode_model_output(model_output1, origin_batch_size=origin_batch_size) + model_output0 = self._create_unpad_decode_model_output( + model_output0, origin_batch_size=origin_batch_size, microbatch_index=0 + ) + model_output1 = self._create_unpad_decode_model_output( + model_output1, origin_batch_size=origin_batch_size, microbatch_index=1 + ) else: model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) @@ -939,8 +1049,12 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state1.init_att_state() model_output0, model_output1 = self._overlap_tpsp_token_forward(infer_state0, infer_state1=infer_state1) - model_output0 = self._create_unpad_decode_model_output(model_output0, origin_batch_size=origin_batch_size) - model_output1 = self._create_unpad_decode_model_output(model_output1, origin_batch_size=origin_batch_size) + model_output0 = self._create_unpad_decode_model_output( + model_output0, origin_batch_size=origin_batch_size, microbatch_index=0 + ) + model_output1 = self._create_unpad_decode_model_output( + model_output1, origin_batch_size=origin_batch_size, microbatch_index=1 + ) return model_output0, model_output1 @@ -986,14 +1100,19 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state model_output = ModelOutput(logits=predict_logits.contiguous(), prompt_logics=infer_state.prompt_logics) model_output1 = ModelOutput(logits=predict_logits1.contiguous(), prompt_logics=infer_state1.prompt_logics) - if self.is_mtp_mode: + if self.spec_adapter is not None: + spec_context = self.spec_adapter.create_forward_context(self, infer_state) input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) - input_embs1 = self.pre_infer._tpsp_allgather(input=input_embs1, infer_state=infer_state1) if infer_state.need_dp_prefill_balance: input_embs = infer_state._all_to_all_unbalance_get(data=input_embs) + spec_context.capture_final_hidden(input_embs.contiguous()) + + if self.spec_adapter is not None: + spec_context1 = self.spec_adapter.create_forward_context(self, infer_state1) + input_embs1 = self.pre_infer._tpsp_allgather(input=input_embs1, infer_state=infer_state1) + if infer_state1.need_dp_prefill_balance: input_embs1 = infer_state1._all_to_all_unbalance_get(data=input_embs1) - model_output.mtp_main_output_hiddens = input_embs.contiguous() - model_output1.mtp_main_output_hiddens = input_embs1.contiguous() + spec_context1.capture_final_hidden(input_embs1.contiguous()) return model_output, model_output1 @@ -1024,11 +1143,15 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: model_output = ModelOutput(logits=predict_logits.contiguous()) model_output1 = ModelOutput(logits=predict_logits1.contiguous()) - if self.is_mtp_mode: + if self.spec_adapter is not None: + spec_context = self.spec_adapter.create_forward_context(self, infer_state) input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) + spec_context.capture_final_hidden(input_embs.contiguous()) + + if self.spec_adapter is not None: + spec_context1 = self.spec_adapter.create_forward_context(self, infer_state1) input_embs1 = self.pre_infer._tpsp_allgather(input=input_embs1, infer_state=infer_state1) - model_output.mtp_main_output_hiddens = input_embs.contiguous() - model_output1.mtp_main_output_hiddens = input_embs1.contiguous() + spec_context1.capture_final_hidden(input_embs1.contiguous()) if infer_state.is_cuda_graph: model_output.to_no_ref_tensor() @@ -1249,3 +1372,10 @@ def _gen_special_model_input(self, token_num: int): special_model_input["mtp_draft_input_hiddens"] = None return special_model_input + + def _gen_mtp_draft_special_model_input(self, token_num: int): + return { + "mtp_draft_input_hiddens": torch.randn( + token_num, self.config["hidden_size"], dtype=self.data_type, device="cuda" + ) + } diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 58922b8055..7520c1aa45 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -1,8 +1,12 @@ import torch -from dataclasses import dataclass, field +from dataclasses import dataclass from typing import Optional from typing import List -from lightllm.utils.envs_utils import enable_diverse_mode_gqa_decode_fast_kernel +from lightllm.utils.envs_utils import ( + enable_diverse_mode_gqa_decode_fast_kernel, + enable_dynamic_mtp_verify, + enable_triton_mtp_kernel, +) from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor @@ -55,6 +59,13 @@ class ModelInput: # mtp_draft_input_hiddens 用于模型 mtp 模式下 # 的 draft 模型的输入 mtp_draft_input_hiddens: Optional[torch.Tensor] = None + # 部分 spec draft 模型会在服务 MTP 模式下执行普通 decode batch + # (例如 Eagle3 commit accepted rows),此时 attention 不应按 MTP 展开布局建参。 + disable_mtp_decode_att: bool = False + # Dynamic verification normally needs arbitrary per-request row groups. + # A full-width plan has the original fixed K+1 layout and can reuse the + # substantially cheaper Static MTP attention parameter construction. + use_static_mtp_layout: bool = False def to_cuda(self): self.check_input() @@ -90,6 +101,12 @@ def to_cuda(self): self.b_shared_seq_len = torch.zeros(size=(batch_size,), dtype=torch.int32, device="cuda") else: self.b_shared_seq_len = self.b_shared_seq_len.cuda(non_blocking=True) + elif not self.is_prefill and (enable_dynamic_mtp_verify() or enable_triton_mtp_kernel()): + batch_size = len(self.b_req_idx) + if self.b_mark_shared_group is None: + self.b_mark_shared_group = torch.ones(size=(batch_size,), dtype=torch.int32, device="cuda") + else: + self.b_mark_shared_group = self.b_mark_shared_group.cuda(non_blocking=True) def __post_init__(self): self.check_input() @@ -108,14 +125,9 @@ class ModelOutput: logits: torch.Tensor # 用于判断 mem_indexes 是否成功写入 req manager 中的事件对象。 prefill_mem_indexes_ready_event: torch.Event = None - - # 专有变量,用于一些特殊的模型,特殊的模式下, 传递一些特殊 - # 的输出变量。只在特殊的模型模式下才会具体使用和生效。 - - # mtp_main_output_hiddens 用于在mtp模式下,llm main model - # 输出最后一层的hidden state 状态用于 draft 模型的 mtp_draft_input_hiddens - # 输入 - mtp_main_output_hiddens: Optional[torch.Tensor] = None + # DSpark dynamic verify 使用的 raw confidence logits,由 draft model + # post layer 产生,proposer 只负责 scatter 到 verify batch。 + mtp_draft_confidence_logits: Optional[torch.Tensor] = None # prompt_logics 用于在开启 return_all_prompt_logics 模式(如 enable_prompt_logprobs)时, # 保存整个 prefill 阶段每一个 token 位置对应的 logits(而非仅最后一个位置的 logits)。 @@ -125,5 +137,5 @@ class ModelOutput: def to_no_ref_tensor(self): self.logits = tensor_to_no_ref_tensor(self.logits) - if self.mtp_main_output_hiddens is not None: - self.mtp_main_output_hiddens = tensor_to_no_ref_tensor(self.mtp_main_output_hiddens) + if self.mtp_draft_confidence_logits is not None: + self.mtp_draft_confidence_logits = tensor_to_no_ref_tensor(self.mtp_draft_confidence_logits) diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 949384b437..c5f588be9f 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -1,16 +1,18 @@ -import os import torch +import torch.distributed as dist import copy import bisect import triton from typing import Optional from lightllm.utils.log_utils import init_logger -from lightllm.utils.envs_utils import get_env_start_args -from lightllm.distributed import dist_group_manager +from lightllm.utils.envs_utils import ( + get_env_start_args, + enable_dynamic_mtp_verify, + get_diverse_max_batch_shared_group_size, +) from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.utils.torch_memory_saver_utils import ( TorchMemorySaverWrapper, - MemoryTag, ) from .infer_struct import InferStateInfo @@ -18,64 +20,160 @@ logger = init_logger(__name__) +def _build_mtp_mark_shared_group_values( + *, + batch_size: int, + mtp_group_size: int, + max_group_size: int, + split_groups: bool, +) -> list: + assert mtp_group_size > 0 + assert max_group_size > 0 + + group_cap = max_group_size if split_groups else mtp_group_size + b_mark_shared_group = [0 for _ in range(batch_size)] + for req_start in range(0, batch_size, mtp_group_size): + req_end = min(req_start + mtp_group_size, batch_size) + for group_start in range(req_start, req_end, group_cap): + group_size = min(group_cap, req_end - group_start) + b_mark_shared_group[group_start + group_size - 1] = group_size + return b_mark_shared_group + + class CudaGraph: # CudaGraph forward pass for the decoding stage. - @staticmethod - def gen_cuda_graph_batch_sizes(max_batch_size=8, tp_world_size: int = 1): - args = get_env_start_args() - mtp_size = args.mtp_step + 1 + def __init__( + self, + max_batch_size=8, + max_len_in_batch=8192, + tp_world_size: int = 1, + graph_split_batch_size: Optional[int] = None, + graph_grow_step_size: Optional[int] = None, + ): + self.graph = {} + self.tp_world_size = tp_world_size + self.mempool = torch.cuda.graph_pool_handle() if torch.cuda.is_available() else None + self.args = get_env_start_args() + self.raw_max_batch_size = max_batch_size + self.max_batch_size = None + self.graph_max_len_in_batch = max_len_in_batch + self.graph_split_batch_size = int( + self.args.graph_split_batch_size if graph_split_batch_size is None else graph_split_batch_size + ) + self.graph_grow_step_size = int( + self.args.graph_grow_step_size if graph_grow_step_size is None else graph_grow_step_size + ) + assert self.graph_split_batch_size > 0 + assert self.graph_grow_step_size > 0 + self.enable_decode_microbatch_overlap = self.args.enable_decode_microbatch_overlap + self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) + self.spec_adapter = None + self.model = None + + self._refresh_cuda_graph_batch_sizes() + return + + def set_spec_adapter(self, spec_adapter, model=None): + self.spec_adapter = spec_adapter + self.model = model + self._refresh_cuda_graph_batch_sizes() + return + def _refresh_cuda_graph_batch_sizes(self): # gen cuda graph batch_sizes # cuda graph gen for batch size = [1, 2, 3, ..., graph_split_batch_size] # and [graph_split_batch_size + graph_grow_step_size, # if the mtp_step is not 0, then the batch_sizes will be multiply of (mtp_step + 1) - graph_split_batch_size = args.graph_split_batch_size * mtp_size - graph_grow_step_size = args.graph_grow_step_size * mtp_size + mtp_step = self._get_decode_graph_mtp_step() + group_size = mtp_step + 1 + self.max_batch_size = (self.raw_max_batch_size // group_size) * group_size + assert self.max_batch_size > 0, "cuda graph max_batch_size must cover at least one decode group" - batch_sizes = [i * mtp_size for i in range(1, args.graph_split_batch_size + 1)] - for _batch_size in range(graph_split_batch_size + graph_grow_step_size, max_batch_size, graph_grow_step_size): + graph_split_batch_size = self.graph_split_batch_size * group_size + graph_grow_step_size = self.graph_grow_step_size * group_size + + batch_sizes = [i * group_size for i in range(1, self.graph_split_batch_size + 1)] + for _batch_size in range( + graph_split_batch_size + graph_grow_step_size, + self.max_batch_size, + graph_grow_step_size, + ): batch_sizes.append(_batch_size) - batch_sizes = list(set([e for e in batch_sizes if e < max_batch_size])) - batch_sizes.append(max_batch_size) + batch_sizes = list(set([e for e in batch_sizes if e < self.max_batch_size])) + batch_sizes.append(self.max_batch_size) batch_sizes.sort() - if args.enable_tpsp_mix_mode: - batch_sizes = [triton.cdiv(e, tp_world_size) * tp_world_size for e in batch_sizes] + if self.args.enable_tpsp_mix_mode: + batch_sizes = [triton.cdiv(e, self.tp_world_size) * self.tp_world_size for e in batch_sizes] batch_sizes = list(set(batch_sizes)) batch_sizes.sort() - assert batch_sizes[-1] == max_batch_size - return batch_sizes - - def __init__(self, max_batch_size=8, max_len_in_batch=8192, tp_world_size: int = 1): - self.graph = {} - self.tp_world_size = tp_world_size - self.mempool = torch.cuda.graph_pool_handle() if torch.cuda.is_available() else None - self.args = get_env_start_args() - self.mtp_step = self.args.mtp_step - self.max_batch_size = max_batch_size - self.graph_max_len_in_batch = max_len_in_batch - self.enable_decode_microbatch_overlap = self.args.enable_decode_microbatch_overlap - self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) - - self.cuda_graph_batch_sizes = self.gen_cuda_graph_batch_sizes( - max_batch_size=max_batch_size, - tp_world_size=tp_world_size, - ) - assert self.cuda_graph_batch_sizes[-1] == self.max_batch_size + self.cuda_graph_batch_sizes = batch_sizes + assert batch_sizes[-1] == self.max_batch_size logger.info(f"cuda graph batch_sizes: {self.cuda_graph_batch_sizes}") + return + + def _get_decode_graph_mtp_step(self) -> int: + if self.spec_adapter is None: + return self.args.mtp_step + return self.spec_adapter.get_decode_graph_mtp_step(self.model) + + def _get_decode_graph_warmup_mtp_step(self) -> int: + if self.spec_adapter is None: + return self.args.mtp_step + get_warmup_step = getattr(self.spec_adapter, "get_decode_graph_warmup_mtp_step", None) + if get_warmup_step is None: + return self.spec_adapter.get_decode_graph_mtp_step(self.model) + return get_warmup_step(self.model) + + def _is_block_draft_model(self) -> bool: + if self.spec_adapter is None or self.model is None: + return False + is_block_draft_model = getattr(self.spec_adapter, "is_block_draft_model", None) + return is_block_draft_model is not None and is_block_draft_model(self.model) def can_run(self, batch_size, max_len_in_batch): return batch_size <= self.max_batch_size and max_len_in_batch <= self.graph_max_len_in_batch - def need_capture(self, batch_size): + def _graph_key( + self, + batch_size: int, + model_context=None, + model_context1=None, + ): + spec_key = None + spec_key1 = None + if self.spec_adapter is not None and model_context is not None: + spec_key = self.spec_adapter.graph_cache_key(model_context, model=self.model) + if self.spec_adapter is not None and model_context1 is not None: + spec_key1 = self.spec_adapter.graph_cache_key(model_context1, model=self.model) + if spec_key is None and spec_key1 is None: + return batch_size + return (batch_size, spec_key, spec_key1) + + def _export_spec_capture(self, infer_state: InferStateInfo): + if self.spec_adapter is None: + return None + return self.spec_adapter.export_graph_capture() + + def _restore_spec_capture(self, infer_state: InferStateInfo, captured_hiddens) -> None: + if self.spec_adapter is not None: + self.spec_adapter.restore_graph_capture(captured_hiddens) + return + + def need_capture( + self, + batch_size, + model_context: Optional[ModelInput] = None, + model_context1: Optional[ModelInput] = None, + ): find_batch_size = self.find_closest_graph_batch_size(batch_size) - if find_batch_size is not None: - return find_batch_size not in self.graph - else: - assert False, "dead code" + return ( + find_batch_size is not None + and self._graph_key(find_batch_size, model_context, model_context1) not in self.graph + ) def find_closest_graph_batch_size(self, batch_size): index = bisect.bisect_left(self.cuda_graph_batch_sizes, batch_size) @@ -85,6 +183,34 @@ def find_closest_graph_batch_size(self, batch_size): else: return None + def _make_warmup_mtp_index(self, batch_size: int) -> torch.Tensor: + mtp_step = self._get_decode_graph_warmup_mtp_step() + if mtp_step <= 0: + return torch.zeros(batch_size, dtype=torch.int32, device="cuda") + return torch.arange(batch_size, dtype=torch.int32, device="cuda") % (mtp_step + 1) + + def _make_warmup_seq_len(self, batch_size: int) -> torch.Tensor: + mtp_step = self._get_decode_graph_warmup_mtp_step() + if mtp_step > 0: + group_size = mtp_step + 1 + return torch.arange(batch_size, dtype=torch.int32, device="cuda") % group_size + 2 + return torch.full((batch_size,), 2, dtype=torch.int32, device="cuda") + + def _make_warmup_mtp_mark_shared_group(self, batch_size: int) -> Optional[torch.Tensor]: + mtp_step = self._get_decode_graph_warmup_mtp_step() + if mtp_step <= 0: + return None + + mtp_group_size = mtp_step + 1 + max_group_size = get_diverse_max_batch_shared_group_size() + b_mark_shared_group = _build_mtp_mark_shared_group_values( + batch_size=batch_size, + mtp_group_size=mtp_group_size, + max_group_size=max_group_size, + split_groups=not self._is_block_draft_model(), + ) + return torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cuda") + def _capture_decode(self, decode_func, infer_state: InferStateInfo): graph_obj = torch.cuda.CUDAGraph() input_ids = infer_state.input_ids @@ -112,10 +238,58 @@ def _capture_decode(self, decode_func, infer_state: InferStateInfo): with self.torch_memory_saver.cuda_graph(graph_obj, pool=self.mempool): model_output = decode_func(infer_state) - self.graph[batch_size] = (graph_obj, infer_state, model_output) + spec_capture = self._export_spec_capture(infer_state) + self.graph[self._graph_key(batch_size, infer_state)] = (graph_obj, infer_state, model_output, spec_capture) graph_obj.replay() + self._record_capture_replay_infer_cost_ms( + graph_obj=graph_obj, + batch_size=batch_size, + is_draft_model=self._is_draft_model_capture(infer_state), + ) return model_output + def _record_capture_replay_infer_cost_ms( + self, + graph_obj: torch.cuda.CUDAGraph, + batch_size: int, + is_draft_model: bool, + ) -> None: + if not enable_dynamic_mtp_verify(): + return + if is_draft_model and self.args.mtp_mode == "dspark": + # DSpark's planner uses target verify cost plus confidence-derived + # capacity estimates. Draft block cost is not part of the decision, + # so avoid adding a runtime barrier/synchronize on lazy draft graph + # capture. + return + + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + dist.barrier(group=dist.group.WORLD) + start_event = torch.cuda.Event(enable_timing=True) + end_event = torch.cuda.Event(enable_timing=True) + graph_obj.replay() + start_event.record() + graph_obj.replay() + end_event.record() + end_event.synchronize() + infer_cost_ms_tensor = torch.tensor([start_event.elapsed_time(end_event)], dtype=torch.float32, device="cuda") + dist.all_reduce(infer_cost_ms_tensor, op=dist.ReduceOp.MIN, group=dist.group.WORLD) + infer_cost_ms = float(infer_cost_ms_tensor.item()) + g_infer_context.record_dynamic_mtp_infer_cost( + batch_size=batch_size, + infer_cost_ms=infer_cost_ms, + is_draft_model=is_draft_model, + ) + return + + def _is_draft_model_capture(self, infer_state: InferStateInfo) -> bool: + if infer_state.mtp_draft_input_hiddens is not None or getattr(infer_state, "is_draft_model", False): + return True + if self.spec_adapter is None or self.model is None: + return False + return self.spec_adapter.is_draft_model(self.model) + def _capture_decode_overlap( self, decode_func, @@ -146,14 +320,23 @@ def _capture_decode_overlap( with self.torch_memory_saver.cuda_graph(graph_obj, pool=self.mempool): model_output, model_output1 = decode_func(infer_state, infer_state1) - self.graph[batch_size] = ( + spec_capture = self._export_spec_capture(infer_state) + spec_capture1 = self._export_spec_capture(infer_state1) + self.graph[self._graph_key(batch_size, infer_state, infer_state1)] = ( graph_obj, infer_state, infer_state1, model_output, model_output1, + spec_capture, + spec_capture1, ) graph_obj.replay() + self._record_capture_replay_infer_cost_ms( + graph_obj=graph_obj, + batch_size=batch_size, + is_draft_model=self._is_draft_model_capture(infer_state), + ) return model_output, model_output1 def capture_decode( @@ -174,9 +357,10 @@ def capture_decode( def _replay(self, infer_state: InferStateInfo): batch_size = infer_state.input_ids.shape[0] - graph_obj, graph_infer_state, graph_output = self.graph[batch_size] + graph_obj, graph_infer_state, graph_output, spec_capture = self.graph[self._graph_key(batch_size, infer_state)] graph_infer_state.copy_for_cuda_graph(infer_state) graph_obj.replay() + self._restore_spec_capture(infer_state, spec_capture) return graph_output def _replay_overlap( @@ -191,10 +375,14 @@ def _replay_overlap( graph_infer_state1, graph_model_output, graph_model_output1, - ) = self.graph[batch_size] + spec_capture, + spec_capture1, + ) = self.graph[self._graph_key(batch_size, infer_state, infer_state1)] graph_infer_state.copy_for_cuda_graph(infer_state) graph_infer_state1.copy_for_cuda_graph(infer_state1) graph_obj.replay() + self._restore_spec_capture(infer_state, spec_capture) + self._restore_spec_capture(infer_state1, spec_capture1) return graph_model_output, graph_model_output1 def replay(self, infer_state, infer_state1=None): @@ -211,20 +399,18 @@ def warmup(self, model): from .basemodel import TpPartBaseModel model: TpPartBaseModel = model - # decode cuda graph init for batch_size in self.cuda_graph_batch_sizes[::-1]: - seq_len = 2 - total_token_num = batch_size * seq_len max_len_in_batch = self.graph_max_len_in_batch input_ids = torch.tensor([1 for _ in range(batch_size)], dtype=torch.int64, device="cuda") mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) - b_seq_len = torch.empty(batch_size, dtype=torch.int32, device="cuda") - b_seq_len.fill_(seq_len) - b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_seq_len = self._make_warmup_seq_len(batch_size) + total_token_num = int(b_seq_len.sum().item()) + b_mtp_index = self._make_warmup_mtp_index(batch_size) + b_mark_shared_group = self._make_warmup_mtp_mark_shared_group(batch_size) model_input = ModelInput( batch_size=batch_size, @@ -236,6 +422,7 @@ def warmup(self, model): b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, + b_mark_shared_group=b_mark_shared_group, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], @@ -243,6 +430,20 @@ def warmup(self, model): ) model_output: ModelOutput = model.forward(model_input) del model_output + if ( + enable_dynamic_mtp_verify() + and self.args.mtp_mode == "eagle3" + and self.spec_adapter is not None + and not self.spec_adapter.is_draft_model(model) + and batch_size % (self.args.mtp_step + 1) == 0 + ): + # Dynamic Eagle3's profitable full-width state reuses the + # fixed K+1 FA3 layout. Capture that graph variant eagerly; + # otherwise its distinct spec key falls back to eager decode + # during the measurement and can be much slower than Static. + model_input.use_static_mtp_layout = True + model_output = model.forward(model_input) + del model_output del input_ids del mem_indexes del b_req_idx @@ -268,22 +469,20 @@ def warmup_overlap(self, model): from .basemodel import TpPartBaseModel model: TpPartBaseModel = model - for batch_size in self.cuda_graph_batch_sizes[::-1]: decode_batches = [] for micro_batch_index in [0, 1]: # dummy decoding, capture the cudagraph - seq_len = 2 - total_token_num = batch_size * seq_len max_len_in_batch = self.graph_max_len_in_batch input_ids = torch.tensor([1 for _ in range(batch_size)], dtype=torch.int64, device="cuda") mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) - b_seq_len = torch.empty(batch_size, dtype=torch.int32, device="cuda") - b_seq_len.fill_(seq_len) - b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_seq_len = self._make_warmup_seq_len(batch_size) + total_token_num = int(b_seq_len.sum().item()) + b_mtp_index = self._make_warmup_mtp_index(batch_size) + b_mark_shared_group = self._make_warmup_mtp_mark_shared_group(batch_size) micro_batch = ModelInput( is_prefill=False, @@ -293,6 +492,7 @@ def warmup_overlap(self, model): max_kv_seq_len=max_len_in_batch, input_ids=input_ids, b_mtp_index=b_mtp_index, + b_mark_shared_group=b_mark_shared_group, mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_seq_len=b_seq_len, diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index 10c35759aa..0ce3c2e0b9 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -1,18 +1,17 @@ import torch -import triton import collections from lightllm.common.kv_cache_mem_manager import MemoryManager -from lightllm.common.req_manager import ReqManager from lightllm.distributed import CustomProcessGroup -from typing import Tuple, Any, Optional, List +from typing import TYPE_CHECKING, Optional, List from .triton_kernel.gen_prefill_params import gen_prefill_params from .triton_kernel.gen_decode_params import gen_decode_params -from .triton_kernel.multimodal_emb import mark_multimodal_obj -from .batch_objs import ModelInput from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.dist_utils import get_global_dp_rank, get_dp_world_size from .attention import BasePrefillAttState, BaseDecodeAttState +if TYPE_CHECKING: + from lightllm.common.req_manager import ReqManager + class InferStateInfo: """ @@ -50,7 +49,7 @@ def __init__(self): self.is_prefill: bool = None self.mem_manager: MemoryManager = None - self.req_manager: ReqManager = None + self.req_manager: "ReqManager" = None self.mem_index: torch.Tensor = None @@ -93,6 +92,7 @@ def __init__(self): # 在开启 mtp_mode 时,mtp draft model # 的输入会用到,其他模型和场景都不会用到 self.mtp_draft_input_hiddens: Optional[torch.Tensor] = None + self.disable_mtp_decode_att: bool = False # 在单节点多dp的运行模式下,在进行prefill的阶段,如果出现了dp之间数据不平衡的现象, # 可以将推理的数据,进行重新分配到各个dp,在做 att 之前,重新 all to all 到各自的 @@ -130,6 +130,19 @@ def init_some_extra_state(self, model): ) = gen_decode_params(self.b_seq_len) self.b_kv_start_loc = self.b1_cu_kv_seq_len[0:-1] + @staticmethod + def build_draft_query_position_ids( + *, + selected_seq_len: torch.Tensor, + b_position_delta: Optional[torch.Tensor], + draft_step: int, + ) -> torch.Tensor: + offsets = torch.arange(draft_step, dtype=torch.long, device=selected_seq_len.device) + position_ids = selected_seq_len.to(dtype=torch.long)[:, None] + offsets[None, :] + if b_position_delta is not None: + position_ids = position_ids + b_position_delta.to(dtype=torch.long)[:, None] + return position_ids + def init_att_state(self): if self.is_prefill: self.prefill_att_state.init_state() diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index 1c1148a55d..a4ccda2ff2 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -1,6 +1,4 @@ -import os import torch -import copy import bisect import triton from typing import List, Tuple @@ -8,7 +6,6 @@ from lightllm.utils.log_utils import init_logger from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor -from lightllm.distributed import dist_group_manager from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from .infer_struct import InferStateInfo from .cuda_graph import CudaGraph @@ -61,10 +58,7 @@ def can_run(self, handle_token_num: int): def need_capture(self, handle_token_num: int): finded_handle_token_num = self.find_closest_graph_handle_token_num(handle_token_num=handle_token_num) - if finded_handle_token_num is not None: - return finded_handle_token_num not in self.graph - else: - assert False, "dead code" + return finded_handle_token_num is not None and finded_handle_token_num not in self.graph def find_closest_graph_handle_token_num(self, handle_token_num: int): index = bisect.bisect_left(self.graph_handle_token_nums, handle_token_num) diff --git a/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/__init__.py b/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/__init__.py new file mode 100644 index 0000000000..f8d0ac0a7e --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/__init__.py @@ -0,0 +1,14 @@ +""" +MTP Diverse Attention Module + +MTP (Multi-Token Prediction) Diverse Attention 的实现。 +""" +from .mtp_diverse_attn import token_decode_attention_mtp_diverse_single_token +from .stage1_single_token import mtp_diverse_stage1_single_token +from .stage2_single_token import mtp_diverse_stage2_single_token + +__all__ = [ + "token_decode_attention_mtp_diverse_single_token", + "mtp_diverse_stage1_single_token", + "mtp_diverse_stage2_single_token", +] diff --git a/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/mtp_diverse_attn.py b/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/mtp_diverse_attn.py new file mode 100644 index 0000000000..101a285e9d --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/mtp_diverse_attn.py @@ -0,0 +1,61 @@ +import torch +from .stage1_single_token import mtp_diverse_stage1_single_token +from .stage2_single_token import mtp_diverse_stage2_single_token +from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size + + +@torch.no_grad() +def token_decode_attention_mtp_diverse_single_token( + q, + k, + v, + Req_to_tokens, + B_req_idx, + b_seq_len, + b_mark_shared_group, + out=None, + alloc_tensor_func=torch.empty, +): + batch_size = b_seq_len.shape[0] + num_heads = q.shape[1] + head_dim = q.shape[2] + + if out is None: + o_tensor = alloc_tensor_func(q.shape, dtype=q.dtype, device=q.device) + else: + o_tensor = out + + max_kv_len = Req_to_tokens.shape[1] + + if batch_size <= 16: + block_num = 128 + elif batch_size <= 64: + block_num = 64 + else: + block_num = 32 + + mid_o = alloc_tensor_func([batch_size, num_heads, block_num, head_dim], dtype=q.dtype, device=q.device) + mid_o_logsumexp = alloc_tensor_func([batch_size, num_heads, block_num], dtype=torch.float32, device=q.device) + + BLOCK_N = mtp_diverse_stage1_single_token( + q=q, + k=k, + v=v, + Req_to_tokens=Req_to_tokens, + B_req_idx=B_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + max_kv_len=max_kv_len, + mid_out=mid_o, + mid_out_logsumexp=mid_o_logsumexp, + block_batch=get_diverse_max_batch_shared_group_size(), + ) + + mtp_diverse_stage2_single_token( + mid_out=mid_o, + mid_out_logsumexp=mid_o_logsumexp, + B_Seqlen=b_seq_len, + out=o_tensor, + block_n=BLOCK_N, + ) + return o_tensor diff --git a/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/stage1_single_token.py b/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/stage1_single_token.py new file mode 100644 index 0000000000..9765d60aaf --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/stage1_single_token.py @@ -0,0 +1,333 @@ +""" +MTP Diverse Attention Stage1 Kernel - Single Token Per Request Mode + +简化版本(参考 int8kv diverse stage1): +- 组内请求 [q1, q2, q3, q4],group mark [0, 0, 0, 4] +- KV slots [kv0, kv1, kv2, kv3, kv4] +- 可见性:q1->[kv0], q2->[kv0,kv1], q3->[kv0,kv1,kv2], q4->[kv0,kv1,kv2,kv3,kv4] + +核心逻辑: +- 只由组内最后一个请求(b_mark_shared_group != 0)触发计算 +- 一次加载组内所有请求的 Q 和 KV +- 每个请求单独做可见性检查(基于各自的 seq_len) +- 中间结果按 kv block 存储,供 Stage2 聚合 +""" +import torch +import triton +import triton.language as tl +from typing import Optional +from lightllm.common.triton_utils.autotuner import autotune, Autotuner +from lightllm.utils.device_utils import is_hopper + + +def get_test_configs(): + configs = [] + for block_n in [16, 32, 64]: + for num_warps in [2, 4, 8]: + for num_stages in [2, 3, 4]: + for warp_specialize in [True, False] if is_hopper() else [False]: + # warp_specialize only support hopper + configs.append( + { + "BLOCK_N": block_n, + "num_warps": num_warps, + "num_stages": num_stages, + "warp_specialize": warp_specialize, + } + ) + return configs + + +def get_static_key(q, k, block_batch): + key_params = { + "gqa_group_size": int(q.shape[1] // k.shape[1]), + "q_head_dim": int(q.shape[2]), + "block_batch": block_batch, + "out_dtype": str(q.dtype), + } + return key_params + + +def get_run_key(q, max_kv_len): + batch_size = q.shape[0] + return batch_size * 1000 * 1000 * 1000 + max_kv_len + + +@triton.jit +def _fwd_kernel_mtp_diverse_stage1_single_token( + Q, + stride_qb, + stride_qh, + stride_qd, + K, + stride_kbs, + stride_kh, + stride_kd, + V, + stride_vbs, + stride_vh, + stride_vd, + sm_scale, + Req_to_tokens, + stride_req_to_tokens_b, + stride_req_to_tokens_s, + B_req_idx, + b_seq_len, + b_mark_shared_group, + Mid_O, + stride_mid_ob, + stride_mid_oh, + stride_mid_os, + stride_mid_od, + Mid_O_LogExpSum, + stride_mid_o_eb, + stride_mid_o_eh, + stride_mid_o_es, + gqa_group_size, + BLOCK_HEAD: tl.constexpr, + BLOCK_BATCH: tl.constexpr, + BLOCK_HEADDIM: tl.constexpr, + BLOCK_N: tl.constexpr, + warp_specialize: tl.constexpr, +): + block_index = tl.program_id(0) + cur_kv_head = tl.program_id(1) + cur_batch = tl.program_id(2) + grid_block_num = tl.num_programs(0) + + shared_batch_group_size = tl.load(b_mark_shared_group + cur_batch) + if shared_batch_group_size == 0: + return + + cur_batch_start = cur_batch - (shared_batch_group_size - 1) + + # ---- batch lane: 不再回卷索引,使用mask ---- + offs_b = tl.arange(0, BLOCK_BATCH) + batch_idx = tl.where(offs_b < shared_batch_group_size, cur_batch_start + offs_b, cur_batch) + + # load seq len + batch_seq_lens = tl.load(b_seq_len + batch_idx) + max_seq_len = tl.max(batch_seq_lens, axis=0) + block_num = tl.cdiv(max_seq_len, BLOCK_N) + if block_index >= block_num: + return + + batch_seq_lens = tl.broadcast_to(batch_seq_lens[:, None], (BLOCK_BATCH, BLOCK_HEAD)) + batch_seq_lens = batch_seq_lens.reshape((BLOCK_BATCH * BLOCK_HEAD,)) + + # ---- head lane: 不再next_pow2回卷,使用mask --- - + offs_h = tl.arange(0, BLOCK_HEAD) + q_head_idx = tl.where(offs_h < gqa_group_size, cur_kv_head * gqa_group_size + offs_h, cur_kv_head * gqa_group_size) + offs_d = tl.arange(0, BLOCK_HEADDIM) + + off_q = batch_idx[:, None, None] * stride_qb + q_head_idx[None, :, None] * stride_qh + offs_d[None, None, :] + q = tl.load(Q + off_q) + q_flat = tl.reshape(q, (BLOCK_BATCH * BLOCK_HEAD, BLOCK_HEADDIM)) + + sum_exp = tl.zeros([BLOCK_BATCH * BLOCK_HEAD], dtype=tl.float32) + max_logic = tl.full([BLOCK_BATCH * BLOCK_HEAD], float("-inf"), dtype=tl.float32) + acc = tl.zeros([BLOCK_BATCH * BLOCK_HEAD, BLOCK_HEADDIM], dtype=tl.float32) + + cur_batch_req_idx = tl.load(B_req_idx + cur_batch) + + for iter_block_index in tl.range(block_index, block_num, grid_block_num, warp_specialize=warp_specialize): + offs_n_new = iter_block_index * BLOCK_N + tl.arange(0, BLOCK_N) + offs_n_refator = tl.where(offs_n_new < max_seq_len, offs_n_new, max_seq_len - 1) + + k_loc = tl.load( + Req_to_tokens + stride_req_to_tokens_b * cur_batch_req_idx + offs_n_refator * stride_req_to_tokens_s, + ).to(tl.int64) + off_k = k_loc[None, :] * stride_kbs + cur_kv_head * stride_kh + offs_d[:, None] + off_v = k_loc[:, None] * stride_vbs + cur_kv_head * stride_vh + offs_d[None, :] + k = tl.load(K + off_k) + v = tl.load(V + off_v) + att = tl.dot(q_flat, k) + att *= sm_scale + att = tl.where(offs_n_new[None, :] < batch_seq_lens[:, None], att, -1000000000.0) + cur_max = tl.max(att, axis=1) + new_max = tl.maximum(cur_max, max_logic) + + exp_logic = tl.exp(att - new_max[:, None]) + logic_scale = tl.exp(max_logic - new_max) + + acc *= logic_scale[:, None] + acc += tl.dot(exp_logic.to(v.dtype), v) + + sum_exp = sum_exp * logic_scale + tl.sum(exp_logic, axis=1) + max_logic = new_max + + mid_o_val = acc / sum_exp[:, None] + mid_lse_val = max_logic + tl.log(sum_exp) + + off_mid_o = ( + batch_idx[:, None, None] * stride_mid_ob + + q_head_idx[None, :, None] * stride_mid_oh + + block_index * stride_mid_os + + offs_d[None, None, :] * stride_mid_od + ) + off_mid_lse = ( + batch_idx[:, None] * stride_mid_o_eb + q_head_idx[None, :] * stride_mid_o_eh + block_index * stride_mid_o_es + ) + + tl.store(Mid_O + off_mid_o, mid_o_val.reshape((BLOCK_BATCH, BLOCK_HEAD, BLOCK_HEADDIM))) + tl.store(Mid_O_LogExpSum + off_mid_lse, mid_lse_val.reshape((BLOCK_BATCH, BLOCK_HEAD))) + + +@autotune( + kernel_name="_fwd_kernel_mtp_diverse_stage1_single_token:v2", + configs_gen_func=get_test_configs, + static_key_func=get_static_key, + run_key_func=get_run_key, + mutates_args=["mid_out", "mid_out_logsumexp"], +) +def mtp_diverse_stage1_single_token( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + Req_to_tokens: torch.Tensor, + B_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + b_mark_shared_group: torch.Tensor, + max_kv_len: int, + mid_out: torch.Tensor, + mid_out_logsumexp: torch.Tensor, + block_batch: int, + run_config: Optional[dict] = None, +): + """ + MTP Diverse Attention Stage1 - Single Token Per Request Mode + + b_seq_len: 每个请求可见的 KV 数量,组内递增 + 例如组内 [q1, q2, q3, q4] 对应 b_seq_len [2, 3, 4, 5] + """ + if not run_config: + run_config = {"BLOCK_N": 16, "num_warps": 2, "num_stages": 2, "warp_specialize": False} # 默认配置 + + BLOCK_N = run_config["BLOCK_N"] + num_warps = run_config["num_warps"] + num_stages = run_config["num_stages"] + warp_specialize = run_config.get("warp_specialize", False) + BLOCK_BATCH = triton.next_power_of_2(block_batch) + + assert q.dim() == 3 and k.dim() == 3 and v.dim() == 3 + assert q.is_cuda and k.is_cuda and v.is_cuda + assert Req_to_tokens.is_cuda and B_req_idx.is_cuda and b_seq_len.is_cuda and b_mark_shared_group.is_cuda + assert mid_out.is_cuda and mid_out_logsumexp.is_cuda + + Lq, Lk = int(q.shape[2]), int(k.shape[2]) + assert Lq == Lk + assert Lk in {16, 32, 64, 128} + batch = int(B_req_idx.shape[0]) + kv_head_num = int(k.shape[1]) + q_head_num = int(q.shape[1]) + assert q_head_num % kv_head_num == 0 + gqa_group_size = q_head_num // kv_head_num + BLOCK_HEAD = triton.next_power_of_2(gqa_group_size) + assert q.stride(-1) == k.stride(-1) == v.stride(-1) == 1 + + grid_num = mid_out.shape[2] + + sm_scale = 1.0 / (Lk ** 0.5) + # 固定 grid(graph friendly) + grid = (grid_num, kv_head_num, batch) + _fwd_kernel_mtp_diverse_stage1_single_token[grid]( + Q=q, + stride_qb=q.stride(0), + stride_qh=q.stride(1), + stride_qd=q.stride(2), + K=k, + stride_kbs=k.stride(0), + stride_kh=k.stride(1), + stride_kd=k.stride(2), + V=v, + stride_vbs=v.stride(0), + stride_vh=v.stride(1), + stride_vd=v.stride(2), + sm_scale=sm_scale, + Req_to_tokens=Req_to_tokens, + stride_req_to_tokens_b=Req_to_tokens.stride(0), + stride_req_to_tokens_s=Req_to_tokens.stride(1), + B_req_idx=B_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + Mid_O=mid_out, + stride_mid_ob=mid_out.stride(0), + stride_mid_oh=mid_out.stride(1), + stride_mid_os=mid_out.stride(2), + stride_mid_od=mid_out.stride(3), + Mid_O_LogExpSum=mid_out_logsumexp, + stride_mid_o_eb=mid_out_logsumexp.stride(0), + stride_mid_o_eh=mid_out_logsumexp.stride(1), + stride_mid_o_es=mid_out_logsumexp.stride(2), + gqa_group_size=gqa_group_size, + BLOCK_HEAD=BLOCK_HEAD, + BLOCK_BATCH=BLOCK_BATCH, + BLOCK_HEADDIM=Lk, + BLOCK_N=BLOCK_N, + num_warps=num_warps, + num_stages=num_stages, + warp_specialize=warp_specialize, + ) + return BLOCK_N + + +if __name__ == "__main__": + from lightllm.utils.envs_utils import get_triton_autotune_level + + if get_triton_autotune_level() != 2: + raise Exception("you need set env LIGHTLLM_TRITON_AUTOTUNE_LEVEL=2 to start program.") + + # static params + q_head_dim = 128 + block_batch = 4 + out_dtype = torch.bfloat16 + + batch_sizes = [1, 8, 16, 32, 64, 128] + decode_lengths = [32, 64, 128, 256, 512, 1024, 2048] + + tp_world_size = 2 + q_head_num = 64 // tp_world_size + k_head_num = 8 // tp_world_size + + gqa_group_size = q_head_num // k_head_num + + Autotuner.start_autotune_warmup() + # autotuing kernel + for batch_size in batch_sizes: + for length in decode_lengths: + # Setup test tensors + q = torch.randn(batch_size, q_head_num, q_head_dim, dtype=out_dtype, device="cuda") + k = torch.randn(batch_size * length, k_head_num, q_head_dim, dtype=out_dtype, device="cuda") + v = torch.randn(batch_size * length, k_head_num, q_head_dim, dtype=out_dtype, device="cuda") + Req_to_tokens = torch.arange(0, batch_size * length, dtype=torch.int32, device="cuda").view( + batch_size, length + ) + B_req_idx = torch.arange(batch_size, dtype=torch.int32, device="cuda") + B_seq_len = torch.full((batch_size,), length, dtype=torch.int32, device="cuda") + b_mark_shared_group = torch.ones(batch_size, dtype=torch.int32, device="cuda") + + if batch_size <= 16: + block_num = 128 + elif batch_size <= 64: + block_num = 64 + else: + block_num = 32 + + mid_out = torch.zeros(batch_size, q_head_num, block_num, q_head_dim, dtype=out_dtype, device="cuda") + mid_out_logsumexp = torch.zeros(batch_size, q_head_num, block_num, dtype=out_dtype, device="cuda") + + mtp_diverse_stage1_single_token( + q=q, + k=k, + v=v, + Req_to_tokens=Req_to_tokens, + B_req_idx=B_req_idx, + b_seq_len=B_seq_len, + b_mark_shared_group=b_mark_shared_group, + max_kv_len=length, + mid_out=mid_out, + mid_out_logsumexp=mid_out_logsumexp, + block_batch=block_batch, + ) + + Autotuner.end_autotune_warmup() diff --git a/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/stage2_single_token.py b/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/stage2_single_token.py new file mode 100644 index 0000000000..4c827ae5d9 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/stage2_single_token.py @@ -0,0 +1,202 @@ +""" +MTP Diverse Attention Stage2 Kernel - Single Token Per Request Mode + +参考 int8kv diverse stage3 的简化实现: +- 每个请求独立聚合自己的中间结果 +- 根据 seq_len 确定需要聚合的 kv block 数量 +- 使用 flash attention reweighting 公式 +""" +import torch +import triton +import triton.language as tl +from typing import Optional +from lightllm.common.triton_utils.autotuner import autotune, Autotuner + + +def get_test_configs(): + configs = [] + for num_warps in [2, 4, 8, 16]: + for num_stages in [1, 2, 4, 5, 6]: + configs.append( + { + "num_warps": num_warps, + "num_stages": num_stages, + } + ) + return configs + + +def get_static_key(mid_out, block_n, out): + key_params = { + "q_head_dim": int(mid_out.shape[-1]), + "block_n": block_n, + "out_dtype": str(out.dtype), + } + return key_params + + +def get_run_key(mid_out): + batch_size_head = mid_out.shape[0] * mid_out.shape[1] + block_num = mid_out.shape[2] + return batch_size_head * 1000 * 1000 * 1000 + block_num + + +@triton.jit +def _fwd_kernel_mtp_diverse_stage2_single_token( + B_Seqlen, + Mid_O, # [batch, head, seq_block_num, head_dim] + Mid_O_LogExpSum, # [batch, head, seq_block_num] + O, # [batch, num_heads, head_dim] + stride_mid_ob, + stride_mid_oh, + stride_mid_os, + stride_mid_od, + stride_mid_o_eb, + stride_mid_o_eh, + stride_mid_o_es, + stride_ob, + stride_oh, + stride_od, + mid_out_block_num, + BLOCK_N: tl.constexpr, + BLOCK_DMODEL: tl.constexpr, + NUM_STAGES: tl.constexpr, +): + """ + MTP Diverse Stage2 Kernel - Single Token Per Request Mode + + 每个请求独立聚合前 seq_len 个 kv block 的中间结果。 + """ + cur_batch = tl.program_id(0) + cur_head = tl.program_id(1) + + cur_batch_seq_len = tl.load(B_Seqlen + cur_batch) + + offs_d = tl.arange(0, BLOCK_DMODEL) + + # 计算需要处理的 kv block 数量 + block_n_size = tl.cdiv(cur_batch_seq_len, BLOCK_N) + block_n_size = tl.minimum(block_n_size, mid_out_block_num) + + # 初始化 accumulator + sum_exp = 0.0 + max_logic = -float("inf") + acc = tl.zeros([BLOCK_DMODEL], dtype=tl.float32) + + for block_idx in tl.range(0, block_n_size, 1, num_stages=NUM_STAGES): + # 加载第 block_idx 个 kv block 的中间结果 + offs_mid_o = cur_batch * stride_mid_ob + cur_head * stride_mid_oh + block_idx * stride_mid_os + offs_d[:] + offs_mid_o_logic = cur_batch * stride_mid_o_eb + cur_head * stride_mid_o_eh + block_idx + + mid_o_val = tl.load(Mid_O + offs_mid_o) + logic_val = tl.load(Mid_O_LogExpSum + offs_mid_o_logic) + + # Flash attention reweighting + new_max_logic = tl.maximum(logic_val, max_logic) + logic_scale = tl.exp(max_logic - new_max_logic) + exp_val = tl.exp(logic_val - new_max_logic) + + acc = acc * logic_scale + exp_val * mid_o_val + sum_exp = sum_exp * logic_scale + exp_val + max_logic = new_max_logic + + # 归一化并存储结果 + offs_o = cur_batch * stride_ob + cur_head * stride_oh + offs_d + tl.store(O + offs_o, acc / sum_exp) + + return + + +@autotune( + kernel_name="_fwd_kernel_mtp_diverse_stage2_single_token:v2", + configs_gen_func=get_test_configs, + static_key_func=get_static_key, + run_key_func=get_run_key, + mutates_args=["out"], +) +@torch.no_grad() +def mtp_diverse_stage2_single_token( + mid_out: torch.Tensor, + mid_out_logsumexp: torch.Tensor, + B_Seqlen: torch.Tensor, + out: torch.Tensor, + block_n: int, + run_config: Optional[dict] = None, +): + if not run_config: + run_config = {"num_warps": 4, "num_stages": 2} + + num_warps = run_config["num_warps"] + num_stages = run_config["num_stages"] + + Lk = mid_out.shape[-1] + assert Lk in {16, 32, 64, 128} + batch, head_num = mid_out.shape[0], mid_out.shape[1] + grid = (batch, head_num) + + _fwd_kernel_mtp_diverse_stage2_single_token[grid]( + B_Seqlen=B_Seqlen, + Mid_O=mid_out, + Mid_O_LogExpSum=mid_out_logsumexp, + O=out, + stride_mid_ob=mid_out.stride(0), + stride_mid_oh=mid_out.stride(1), + stride_mid_os=mid_out.stride(2), + stride_mid_od=mid_out.stride(3), + stride_mid_o_eb=mid_out_logsumexp.stride(0), + stride_mid_o_eh=mid_out_logsumexp.stride(1), + stride_mid_o_es=mid_out_logsumexp.stride(2), + stride_ob=out.stride(0), + stride_oh=out.stride(1), + stride_od=out.stride(2), + mid_out_block_num=mid_out.shape[2], + BLOCK_N=block_n, + BLOCK_DMODEL=Lk, + NUM_STAGES=num_stages, + num_warps=num_warps, + num_stages=num_stages, + ) + return + + +if __name__ == "__main__": + from lightllm.utils.envs_utils import get_triton_autotune_level + + if get_triton_autotune_level() != 2: + raise Exception("you need set env LIGHTLLM_TRITON_AUTOTUNE_LEVEL=2 to start program.") + + q_head_dim = 128 + tp_world_size = 2 + out_dtype = torch.float + + batch_sizes = [1, 8, 16, 32, 64, 128] + q_head_num = 64 // tp_world_size + + Autotuner.start_autotune_warmup() + # autotuing kernel + for batch_size in batch_sizes: + for block_n in [16, 32, 64, 128]: + for out_dtype in [torch.float16, torch.bfloat16]: + if batch_size <= 16: + block_num = 128 + elif batch_size <= 64: + block_num = 64 + else: + block_num = 32 + + mid_out = torch.randn( + batch_size, q_head_num, block_num, q_head_dim, dtype=torch.bfloat16, device="cuda" + ) + mid_out_logsumexp = torch.randn(batch_size, q_head_num, block_num, dtype=torch.float32, device="cuda") + B_Seqlen = torch.full((batch_size,), 8196, dtype=torch.int32, device="cuda") + out = torch.zeros(batch_size, q_head_num, q_head_dim, dtype=out_dtype, device="cuda") + + mtp_diverse_stage2_single_token( + mid_out=mid_out, + mid_out_logsumexp=mid_out_logsumexp, + B_Seqlen=B_Seqlen, + out=out, + block_n=block_n, + ) + + Autotuner.end_autotune_warmup() diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py new file mode 100644 index 0000000000..64db53d4df --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -0,0 +1,180 @@ +import triton +import triton.language as tl +from triton.language.standard import _log2, sum, zeros_like +import torch + + +@triton.jit +def _fwd_kernel_cumprod_probs( + req_to_next_token_probs, + req_to_next_token_probs_stride, + b_req_idx, + mtp_step, + BLOCK_SIZE: tl.constexpr, +): + cur_index = tl.program_id(0) + cur_req_idx = tl.load(b_req_idx + cur_index * (mtp_step + 1)) + base_ptr = req_to_next_token_probs + cur_req_idx * req_to_next_token_probs_stride + tl.store(base_ptr, 1.0) + + offset = tl.arange(0, BLOCK_SIZE) + store_mask = offset < (mtp_step + 1) + + probs = tl.load(base_ptr + offset, mask=store_mask, other=0.0) + # offset 0 是 target sample,本轮恒接受;只有 draft 条件接受概率需要 clamp。 + probs = tl.where(offset == 0, 1.0, probs) + # 对于 draft probs 中大于 0.99 的值,设置为 0.99,避免错误的值,照成后续的采样操作失败。 + # 对于 draft probs 中小于 0.01 的值,设置为 0.01,避免错误的值,照成后续的采样操作失败。 + # 这样修改后,我们可以做到,对于一个req的所有mtp step 步的概率的cumprod是降序的 + # 从而在采样时,我们只需要进行排序,则选中的位置,必然满足先后关系,避免一些复杂的 + # 额外操作。 + probs = tl.where((offset != 0) & (probs >= 0.99), 0.99, probs) + probs = tl.where((offset != 0) & (probs <= 0.01), 0.01, probs) + + cum_probs = tl.cumprod(probs, axis=0) + + tl.store(base_ptr + offset, cum_probs, mask=store_mask) + return + + +@triton.jit +def _compare_and_swap(x, ids, flip, i: tl.core.constexpr, n_dims: tl.core.constexpr): + n_outer: tl.core.constexpr = x.numel >> n_dims + shape: tl.core.constexpr = [n_outer * 2 ** i, 2, 2 ** (n_dims - i - 1)] + y = tl.core.reshape(x, shape) + # slice left/right with 'stride' 2**(n_dims - i - 1) + mask = tl.core.arange(0, 2)[None, :, None] + left = tl.core.broadcast_to(sum(y * (1 - mask), 1)[:, None, :], shape) + right = tl.core.broadcast_to(sum(y * mask, 1)[:, None, :], shape) + left = tl.core.reshape(left, x.shape) + right = tl.core.reshape(right, x.shape) + + y_idx = tl.core.reshape(ids, shape) + left_idx = tl.core.broadcast_to(sum(y_idx * (1 - mask), 1)[:, None, :], shape) + right_idx = tl.core.broadcast_to(sum(y_idx * mask, 1)[:, None, :], shape) + left_idx = tl.core.reshape(left_idx, x.shape) + right_idx = tl.core.reshape(right_idx, x.shape) + + idtype = tl.core.get_int_dtype(bitwidth=x.dtype.primitive_bitwidth, signed=True) + ileft = left.to(idtype, bitcast=True) + iright = right.to(idtype, bitcast=True) + ix = x.to(idtype, bitcast=True) + + cond = (left > right) != (flip != 0) + + ret = ix ^ tl.core.where(cond, ileft ^ iright, zeros_like(ix)) + new_ids = ids ^ tl.core.where(cond, left_idx ^ right_idx, zeros_like(ids)) + + return ret.to(x.dtype, bitcast=True), new_ids + + +@triton.jit +def _bitonic_merge(x, ids, stage: tl.core.constexpr, order: tl.core.constexpr, n_dims: tl.core.constexpr): + """ + order_type 0 == ascending + order_type 1 == descending + order_type 2 == alternating + """ + n_outer: tl.core.constexpr = x.numel >> n_dims + tl.core.static_assert(stage <= n_dims) + if order == 2: + shape: tl.core.constexpr = [n_outer * 2 ** (n_dims - 1 - stage), 2, 2 ** stage] + flip = tl.core.reshape(tl.core.broadcast_to(tl.core.arange(0, 2)[None, :, None], shape), x.shape) + else: + flip = order + for i in tl.core.static_range(stage): + x, ids = _compare_and_swap(x, ids, flip, i + (n_dims - stage), n_dims) + return x, ids + + +@triton.jit +def argsort(x, ids, dim: tl.core.constexpr = None, descending: tl.core.constexpr = tl.core.CONSTEXPR_0): + _dim: tl.core.constexpr = len(x.shape) - 1 if dim is None else dim + tl.core.static_assert(_dim == len(x.shape) - 1, "only minor dimension is currently supported") + n_dims: tl.core.constexpr = _log2(x.shape[_dim]) + + for i in tl.core.static_range(1, n_dims + 1): + x, ids = _bitonic_merge(x, ids, i, 2 if i < n_dims else descending, n_dims) + return x, ids + + +@triton.jit +def _fwd_kernel_sample_dynamic_mtp_steps( + req_to_next_token_probs, + req_to_next_token_probs_stride, + select_run_reqs, + b_req_idx, + mtp_step, + verify_step, + req_num, + dynamic_batch_size, + BLOCK_SIZE: tl.constexpr, +): + all_num = req_num * (verify_step + 1) + offset = tl.arange(0, BLOCK_SIZE) + mask = offset < all_num + req_offset = offset // (verify_step + 1) + next_token_offset = offset % (verify_step + 1) + original_offset = req_offset * (mtp_step + 1) + next_token_offset + + req_idx_index = tl.load(b_req_idx + original_offset, mask=mask, other=0) + + probs = tl.load( + req_to_next_token_probs + req_idx_index * req_to_next_token_probs_stride + next_token_offset, + mask=mask, + other=-1.0, + ) + + sorted_probs, sorted_ids = argsort(probs, original_offset, descending=True) + + tl.store(select_run_reqs + sorted_ids, 1, mask=mask & (offset < dynamic_batch_size)) + return + + +def sample_dynamic_mtp_req_mask( + dynamic_batch_size: int, + b_req_idx: torch.Tensor, + req_to_next_token_probs: torch.Tensor, + mtp_step: int, + verify_step: int = None, +) -> torch.Tensor: + dynamic_batch_size = int(dynamic_batch_size) + mtp_step = int(mtp_step) + verify_step = mtp_step if verify_step is None else int(verify_step) + assert 0 <= verify_step <= mtp_step + assert b_req_idx.shape[0] % (mtp_step + 1) == 0 + assert req_to_next_token_probs.is_cuda + assert dynamic_batch_size <= b_req_idx.shape[0] + req_num = len(b_req_idx) // (mtp_step + 1) + valid_row_num = req_num * (verify_step + 1) + assert dynamic_batch_size <= valid_row_num + + # cumprod probs for each request + _fwd_kernel_cumprod_probs[(req_num,)]( + req_to_next_token_probs=req_to_next_token_probs, + req_to_next_token_probs_stride=req_to_next_token_probs.stride(0), + b_req_idx=b_req_idx, + mtp_step=mtp_step, + BLOCK_SIZE=triton.next_power_of_2(mtp_step + 1), + num_warps=1, + num_stages=1, + ) + + # 1 为选中, 0 为未选中 + select_run_reqs = torch.zeros((len(b_req_idx),), dtype=torch.int32, device="cuda") + + grid = (1,) + _fwd_kernel_sample_dynamic_mtp_steps[grid]( + req_to_next_token_probs=req_to_next_token_probs, + req_to_next_token_probs_stride=req_to_next_token_probs.stride(0), + select_run_reqs=select_run_reqs, + b_req_idx=b_req_idx, + mtp_step=mtp_step, + verify_step=verify_step, + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + BLOCK_SIZE=triton.next_power_of_2(valid_row_num), + num_warps=1, + num_stages=1, + ) + return select_run_reqs diff --git a/lightllm/common/basemodel/triton_kernel/fa3_utils.py b/lightllm/common/basemodel/triton_kernel/fa3_utils.py index 0a524b63b6..c20d2443d5 100644 --- a/lightllm/common/basemodel/triton_kernel/fa3_utils.py +++ b/lightllm/common/basemodel/triton_kernel/fa3_utils.py @@ -1,7 +1,12 @@ +import torch import triton import triton.language as tl +_DYNAMIC_MTP_FA3_FAST_PATH_MAX_BATCH_SIZE = 1024 +_DYNAMIC_MTP_FA3_COMPACT_BLOCK_SIZE = 256 + + @triton.jit def page_table_copy_kernel( page_table_ptr, @@ -57,34 +62,160 @@ def page_table_copy( ) -def test_page_table_copy(): - import torch +@triton.jit +def _build_dynamic_mtp_fa3_decode_params_kernel( + b_req_idx, + b_seq_len, + b_mark_shared_group, + out_b_q_seq_len, + out_b_kv_seq_len, + out_b_att_req_idx, + out_b_att_seq_len, + batch_size, + hold_req_id, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size - batch_size, seq_len = 2, 8 + mark = tl.load(b_mark_shared_group + offsets, mask=mask, other=0) + is_group_end = mask & (mark > 0) + dst_pos = tl.cumsum(tl.where(is_group_end, 1, 0), axis=0) - 1 - req_to_token_indexs = torch.arange(batch_size * seq_len, dtype=torch.int32).reshape(batch_size, seq_len).cuda() + tl.store(out_b_q_seq_len + offsets, 0, mask=mask) + tl.store(out_b_kv_seq_len + offsets, 0, mask=mask) + tl.store(out_b_att_req_idx + offsets, hold_req_id, mask=mask) + tl.store(out_b_att_seq_len + offsets, 0, mask=mask) - page_table = torch.full((batch_size, seq_len), -1, dtype=torch.int32, device="cuda") + cur_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=hold_req_id) + cur_seq_len = tl.load(b_seq_len + offsets, mask=mask, other=0) - b_req_idx = torch.tensor([0, 2, 1, 3], dtype=torch.int32, device="cuda")[::2] - print(b_req_idx.stride()) + tl.store(out_b_q_seq_len + dst_pos, mark, mask=is_group_end) + tl.store(out_b_kv_seq_len + dst_pos, cur_seq_len, mask=is_group_end) + tl.store(out_b_att_req_idx + dst_pos, cur_req_idx, mask=is_group_end) + tl.store(out_b_att_seq_len + dst_pos, cur_seq_len, mask=is_group_end) - page_table_copy(page_table, req_to_token_indexs, b_req_idx) - print("req_to_token_indexs:") - print(req_to_token_indexs.cpu().numpy()) - print("b_req_idx:", b_req_idx.cpu().numpy()) - print("page_table:") - print(page_table.cpu().numpy()) +@triton.jit +def _count_dynamic_mtp_fa3_decode_params_kernel( + b_mark_shared_group, + out_block_counts, + batch_size, + BLOCK_SIZE: tl.constexpr, +): + block_id = tl.program_id(axis=0) + offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size - for batch in range(batch_size): - src_idx = b_req_idx[batch].item() - expected = req_to_token_indexs[src_idx].cpu().numpy() - got = page_table[batch].cpu().numpy() - assert (expected == got).all(), f"Batch {batch} mismatch: expected {expected}, got {got}" + mark = tl.load(b_mark_shared_group + offsets, mask=mask, other=0) + block_count = tl.sum(tl.where(mask & (mark > 0), 1, 0), axis=0) + tl.store(out_block_counts + block_id, block_count) - print("✅ Test passed!") +@triton.jit +def _compact_dynamic_mtp_fa3_decode_params_kernel( + b_req_idx, + b_seq_len, + b_mark_shared_group, + block_counts, + block_offsets, + out_b_q_seq_len, + out_b_kv_seq_len, + out_b_att_req_idx, + out_b_att_seq_len, + batch_size, + hold_req_id, + BLOCK_SIZE: tl.constexpr, +): + block_id = tl.program_id(axis=0) + offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size + + mark = tl.load(b_mark_shared_group + offsets, mask=mask, other=0) + is_group_end = mask & (mark > 0) + local_pos = tl.cumsum(tl.where(is_group_end, 1, 0), axis=0) - 1 + block_count = tl.load(block_counts + block_id) + block_end_offset = tl.load(block_offsets + block_id) + block_start_offset = block_end_offset - block_count + dst_pos = block_start_offset + local_pos + + cur_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=hold_req_id) + cur_seq_len = tl.load(b_seq_len + offsets, mask=mask, other=0) + + tl.store(out_b_q_seq_len + dst_pos, mark, mask=is_group_end) + tl.store(out_b_kv_seq_len + dst_pos, cur_seq_len, mask=is_group_end) + tl.store(out_b_att_req_idx + dst_pos, cur_req_idx, mask=is_group_end) + tl.store(out_b_att_seq_len + dst_pos, cur_seq_len, mask=is_group_end) + + +@torch.no_grad() +def build_dynamic_mtp_fa3_decode_params( + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + b_mark_shared_group: torch.Tensor, + att_batch_size: int, + hold_req_id: int, +): + assert b_req_idx.is_cuda and b_seq_len.is_cuda and b_mark_shared_group.is_cuda + assert b_req_idx.shape == b_seq_len.shape == b_mark_shared_group.shape + assert b_req_idx.shape[0] == att_batch_size + assert att_batch_size > 0 + + if att_batch_size <= _DYNAMIC_MTP_FA3_FAST_PATH_MAX_BATCH_SIZE: + b_q_seq_len = torch.empty((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) + b_kv_seq_len = torch.empty((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) + b_att_req_idx = torch.empty((att_batch_size,), dtype=torch.int32, device=b_req_idx.device) + b_att_seq_len = torch.empty((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) + + _build_dynamic_mtp_fa3_decode_params_kernel[(1,)]( + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + out_b_q_seq_len=b_q_seq_len, + out_b_kv_seq_len=b_kv_seq_len, + out_b_att_req_idx=b_att_req_idx, + out_b_att_seq_len=b_att_seq_len, + batch_size=att_batch_size, + hold_req_id=hold_req_id, + BLOCK_SIZE=triton.next_power_of_2(att_batch_size), + num_warps=8, + num_stages=1, + ) + return b_q_seq_len, b_kv_seq_len, b_att_req_idx, b_att_seq_len + + b_q_seq_len = torch.zeros((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) + b_kv_seq_len = torch.zeros((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) + b_att_req_idx = torch.full((att_batch_size,), hold_req_id, dtype=torch.int32, device=b_req_idx.device) + b_att_seq_len = torch.zeros((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) + + block_size = _DYNAMIC_MTP_FA3_COMPACT_BLOCK_SIZE + grid = (triton.cdiv(att_batch_size, block_size),) + block_counts = torch.empty((grid[0],), dtype=torch.int32, device=b_mark_shared_group.device) + + _count_dynamic_mtp_fa3_decode_params_kernel[grid]( + b_mark_shared_group=b_mark_shared_group, + out_block_counts=block_counts, + batch_size=att_batch_size, + BLOCK_SIZE=block_size, + num_warps=8, + num_stages=1, + ) + block_offsets = torch.cumsum(block_counts, dim=0, dtype=torch.int32) -if __name__ == "__main__": - test_page_table_copy() + _compact_dynamic_mtp_fa3_decode_params_kernel[grid]( + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + block_counts=block_counts, + block_offsets=block_offsets, + out_b_q_seq_len=b_q_seq_len, + out_b_kv_seq_len=b_kv_seq_len, + out_b_att_req_idx=b_att_req_idx, + out_b_att_seq_len=b_att_seq_len, + batch_size=att_batch_size, + hold_req_id=hold_req_id, + BLOCK_SIZE=block_size, + num_warps=8, + num_stages=1, + ) + return b_q_seq_len, b_kv_seq_len, b_att_req_idx, b_att_seq_len diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index 26e1468bd4..c06337033e 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -1,7 +1,12 @@ +from typing import Optional import triton import triton.language as tl import torch +from lightllm.common.basemodel.batch_objs import ModelInput +from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import sample_dynamic_mtp_req_mask +from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size, get_env_start_args + @triton.jit def _fwd_kernel_mtp_verify( @@ -93,19 +98,34 @@ def _fwd_kernel_mtp_scatter_next_token_ids( req_to_next_token_ids_stride, all_next_token_ids, all_next_token_ids_stride, + req_to_next_token_probs, + req_to_next_token_probs_stride, + all_next_token_probs, + all_next_token_probs_stride, mtp_accept_len, b_req_mtp_start_loc, b_req_idx, mtp_step, + HAS_HAS_NEXT_TOKEN_PROBS: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): - cur_index = tl.program_id(0) req_start_loc = tl.load(b_req_mtp_start_loc + cur_index) accept_len = tl.load(mtp_accept_len + cur_index) cur_req_idx = tl.load(b_req_idx + req_start_loc) offset = tl.arange(0, BLOCK_SIZE) + if HAS_HAS_NEXT_TOKEN_PROBS: + cur_next_token_probs = tl.load( + all_next_token_probs + (req_start_loc + accept_len - 1) * all_next_token_probs_stride + offset, + mask=offset < mtp_step, + other=0.0, + ) + tl.store( + req_to_next_token_probs + cur_req_idx * req_to_next_token_probs_stride + offset, + cur_next_token_probs, + mask=offset < mtp_step, + ) scatter_next_token_ids = tl.load( all_next_token_ids + (req_start_loc + accept_len - 1) * all_next_token_ids_stride + offset, mask=offset < mtp_step, @@ -125,12 +145,31 @@ def mtp_scatter_next_token_ids( all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, mtp_accept_len: torch.Tensor, + req_to_next_token_probs: Optional[torch.Tensor] = None, + all_next_token_probs: Optional[torch.Tensor] = None, ): max_mtp_step = req_to_next_token_ids.shape[1] BLOCK_SIZE = 16 assert max_mtp_step <= BLOCK_SIZE, f"max_mtp_step must be less than {BLOCK_SIZE}" num_reqs = b_req_mtp_start_loc.shape[0] mtp_step = all_next_token_ids.shape[1] + if req_to_next_token_probs is not None: + assert all_next_token_probs is not None + assert all_next_token_probs.shape == all_next_token_ids.shape + + HAS_HAS_NEXT_TOKEN_PROBS = req_to_next_token_probs is not None + # Triton launch 参数阶段不能直接传 None;不开动态 MTP 时这里传一个不会被实际使用的 dummy tensor 即可。 + req_to_next_token_probs_arg = ( + req_to_next_token_probs if req_to_next_token_probs is not None else req_to_next_token_ids + ) + req_to_next_token_probs_stride = ( + req_to_next_token_probs.stride(0) if req_to_next_token_probs is not None else req_to_next_token_ids.stride(0) + ) + all_next_token_probs_arg = all_next_token_probs if all_next_token_probs is not None else all_next_token_ids + all_next_token_probs_stride = ( + all_next_token_probs.stride(0) if all_next_token_probs is not None else all_next_token_ids.stride(0) + ) + grid = (num_reqs,) num_warps = 1 _fwd_kernel_mtp_scatter_next_token_ids[grid]( @@ -138,22 +177,384 @@ def mtp_scatter_next_token_ids( req_to_next_token_ids_stride=req_to_next_token_ids.stride(0), all_next_token_ids=all_next_token_ids, all_next_token_ids_stride=all_next_token_ids.stride(0), + req_to_next_token_probs=req_to_next_token_probs_arg, + req_to_next_token_probs_stride=req_to_next_token_probs_stride, + all_next_token_probs=all_next_token_probs_arg, + all_next_token_probs_stride=all_next_token_probs_stride, mtp_accept_len=mtp_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, mtp_step=mtp_step, + HAS_HAS_NEXT_TOKEN_PROBS=HAS_HAS_NEXT_TOKEN_PROBS, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps, num_stages=1, ) +@triton.jit +def _fwd_kernel_trim_dynamic_mtp_model_input( + input_ids, + out_input_ids, + b_req_idx, + out_b_req_idx, + b_mtp_index, + out_b_mtp_index, + b_seq_len, + out_b_seq_len, + b_position_delta, + out_b_position_delta, + mem_indexes, + out_mem_indexes, + b_shared_seq_len, + out_b_shared_seq_len, + selected_mask, + selected_dst_pos, + batch_size, + HAS_INPUT_IDS: tl.constexpr, + HAS_B_POSITION_DELTA: tl.constexpr, + HAS_MEM_INDEXES: tl.constexpr, + HAS_B_SHARED_SEQ_LEN: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size + selected_i32 = tl.load(selected_mask + offsets, mask=mask, other=0) + selected = selected_i32 != 0 + dst_pos = tl.cumsum(selected_i32, axis=0) - 1 + write_mask = mask & selected + + cur_b_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=0) + cur_b_mtp_index = tl.load(b_mtp_index + offsets, mask=mask, other=0) + cur_b_seq_len = tl.load(b_seq_len + offsets, mask=mask, other=0) + + tl.store(selected_dst_pos + offsets, dst_pos, mask=mask) + tl.store(out_b_req_idx + dst_pos, cur_b_req_idx, mask=write_mask) + tl.store(out_b_mtp_index + dst_pos, cur_b_mtp_index, mask=write_mask) + tl.store(out_b_seq_len + dst_pos, cur_b_seq_len, mask=write_mask) + + if HAS_INPUT_IDS: + input_id = tl.load(input_ids + offsets, mask=mask, other=0) + tl.store(out_input_ids + dst_pos, input_id, mask=write_mask) + + if HAS_B_POSITION_DELTA: + position_delta = tl.load(b_position_delta + offsets, mask=mask, other=0) + tl.store(out_b_position_delta + dst_pos, position_delta, mask=write_mask) + + if HAS_MEM_INDEXES: + mem_index = tl.load(mem_indexes + offsets, mask=mask, other=0) + tl.store(out_mem_indexes + dst_pos, mem_index, mask=write_mask) + + if HAS_B_SHARED_SEQ_LEN: + shared_seq_len = tl.load(b_shared_seq_len + offsets, mask=mask, other=0) + tl.store(out_b_shared_seq_len + dst_pos, shared_seq_len, mask=write_mask) + + return + + +@triton.jit +def _fwd_kernel_rebuild_trimmed_mtp_b_mark_shared_group( + b_req_idx, + out_b_mark_shared_group, + batch_size, + max_batch_shared_group_size: tl.constexpr, + MAX_RUN_SCAN: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size + cur_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=-1) + + prev_same_count = tl.full((BLOCK_SIZE,), 0, tl.int32) + for scan_offset in tl.static_range(1, MAX_RUN_SCAN + 1): + prev_offsets = offsets - scan_offset + prev_mask = mask & (prev_offsets >= 0) + prev_req_idx = tl.load(b_req_idx + prev_offsets, mask=prev_mask, other=-2) + prev_same_count += tl.where(prev_mask & (prev_req_idx == cur_req_idx), 1, 0) + + next_offsets = offsets + 1 + next_req_idx = tl.load(b_req_idx + next_offsets, mask=next_offsets < batch_size, other=-2) + group_pos = prev_same_count % max_batch_shared_group_size + is_group_end = mask & ( + (next_offsets == batch_size) | (next_req_idx != cur_req_idx) | (group_pos == max_batch_shared_group_size - 1) + ) + mark_value = tl.where(is_group_end, group_pos + 1, 0) + tl.store(out_b_mark_shared_group + offsets, mark_value, mask=mask) + return + + +@triton.jit +def _fwd_kernel_pack_selected_rows_2d( + src, + src_stride_0, + src_stride_1, + dst, + dst_stride_0, + dst_stride_1, + selected_mask, + selected_dst_pos, + batch_size, + hidden_size, + BLOCK_N: tl.constexpr, +): + row_id = tl.program_id(0) + col_block_id = tl.program_id(1) + col_offsets = col_block_id * BLOCK_N + tl.arange(0, BLOCK_N) + row_mask = row_id < batch_size + + selected_val = tl.load(selected_mask + row_id, mask=row_mask, other=0) + dst_row = tl.load(selected_dst_pos + row_id, mask=row_mask, other=0) + write_row = row_mask & (selected_val != 0) + col_mask = col_offsets < hidden_size + src_ptrs = src + row_id * src_stride_0 + col_offsets * src_stride_1 + dst_ptrs = dst + dst_row * dst_stride_0 + col_offsets * dst_stride_1 + vals = tl.load(src_ptrs, mask=write_row & col_mask, other=0) + tl.store(dst_ptrs, vals, mask=write_row & col_mask) + + +def _pack_selected_rows_2d( + src: Optional[torch.Tensor], + selected_mask_gpu: torch.Tensor, + selected_dst_pos: torch.Tensor, + dynamic_batch_size: int, +): + if src is None: + return None + + assert src.is_cuda + assert src.ndim == 2 + assert selected_mask_gpu.is_cuda + assert selected_dst_pos.is_cuda + assert src.shape[0] == selected_mask_gpu.shape[0] + + selected_mask_gpu = selected_mask_gpu.to(torch.int32) + hidden_size = src.shape[1] + dst = torch.empty((dynamic_batch_size, hidden_size), dtype=src.dtype, device=src.device) + grid = (src.shape[0], triton.cdiv(hidden_size, 128)) + _fwd_kernel_pack_selected_rows_2d[grid]( + src=src, + src_stride_0=src.stride(0), + src_stride_1=src.stride(1), + dst=dst, + dst_stride_0=dst.stride(0), + dst_stride_1=dst.stride(1), + selected_mask=selected_mask_gpu, + selected_dst_pos=selected_dst_pos, + batch_size=src.shape[0], + hidden_size=hidden_size, + BLOCK_N=128, + num_warps=4, + num_stages=1, + ) + return dst + + +def _rebuild_trimmed_mtp_b_mark_shared_group_from_b_req_idx(b_req_idx: torch.Tensor) -> torch.Tensor: + assert b_req_idx.is_cuda + batch_size = b_req_idx.shape[0] + max_batch_shared_group_size = int(get_diverse_max_batch_shared_group_size()) + assert max_batch_shared_group_size > 0 + if batch_size == 0: + return torch.empty((0,), dtype=torch.int32, device=b_req_idx.device) + + b_mark_shared_group = torch.empty((batch_size,), dtype=torch.int32, device=b_req_idx.device) + BLOCK_SIZE = triton.next_power_of_2(batch_size) + _fwd_kernel_rebuild_trimmed_mtp_b_mark_shared_group[(1,)]( + b_req_idx=b_req_idx, + out_b_mark_shared_group=b_mark_shared_group, + batch_size=batch_size, + max_batch_shared_group_size=max_batch_shared_group_size, + MAX_RUN_SCAN=16, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=8, + num_stages=1, + ) + return b_mark_shared_group + + +def _trim_decode_model_input_inplace( + model_input: ModelInput, + selected_mask_gpu: torch.Tensor, + dynamic_batch_size: int, +) -> ModelInput: + assert not model_input.is_prefill + assert selected_mask_gpu.is_cuda + assert model_input.b_req_idx.is_cuda + assert model_input.b_mtp_index.is_cuda + assert model_input.b_seq_len.is_cuda + + # 动态 MTP 采样阶段已经保证 selected_mask_gpu 恰好选出 dynamic_batch_size 个位置。 + selected_mask_gpu = selected_mask_gpu.to(torch.int32) + old_batch_size = model_input.b_req_idx.shape[0] + selected_dst_pos = torch.empty((old_batch_size,), dtype=torch.int32, device=model_input.b_req_idx.device) + + out_input_ids = None + if model_input.input_ids is not None: + assert model_input.input_ids.is_cuda + out_input_ids = torch.empty( + (dynamic_batch_size,), dtype=model_input.input_ids.dtype, device=model_input.input_ids.device + ) + + out_b_shared_seq_len = None + if model_input.b_shared_seq_len is not None: + assert model_input.b_shared_seq_len.is_cuda + out_b_shared_seq_len = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_shared_seq_len.dtype, device=model_input.b_shared_seq_len.device + ) + + out_b_req_idx = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_req_idx.dtype, device=model_input.b_req_idx.device + ) + out_b_mtp_index = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_mtp_index.dtype, device=model_input.b_mtp_index.device + ) + out_b_seq_len = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_seq_len.dtype, device=model_input.b_seq_len.device + ) + out_b_position_delta = None + if model_input.b_position_delta is not None: + assert model_input.b_position_delta.is_cuda + out_b_position_delta = torch.empty( + (dynamic_batch_size,), + dtype=model_input.b_position_delta.dtype, + device=model_input.b_position_delta.device, + ) + + out_mem_indexes = None + if model_input.mem_indexes is not None: + assert model_input.mem_indexes.is_cuda + out_mem_indexes = torch.empty( + (dynamic_batch_size,), + dtype=model_input.mem_indexes.dtype, + device=model_input.mem_indexes.device, + ) + + dummy_1d = model_input.b_req_idx + BLOCK_SIZE = triton.next_power_of_2(old_batch_size) + grid = (1,) + _fwd_kernel_trim_dynamic_mtp_model_input[grid]( + input_ids=model_input.input_ids if model_input.input_ids is not None else dummy_1d, + out_input_ids=out_input_ids if out_input_ids is not None else dummy_1d, + b_req_idx=model_input.b_req_idx, + out_b_req_idx=out_b_req_idx, + b_mtp_index=model_input.b_mtp_index, + out_b_mtp_index=out_b_mtp_index, + b_seq_len=model_input.b_seq_len, + out_b_seq_len=out_b_seq_len, + b_position_delta=model_input.b_position_delta if model_input.b_position_delta is not None else dummy_1d, + out_b_position_delta=out_b_position_delta if out_b_position_delta is not None else dummy_1d, + mem_indexes=model_input.mem_indexes if model_input.mem_indexes is not None else dummy_1d, + out_mem_indexes=out_mem_indexes if out_mem_indexes is not None else dummy_1d, + b_shared_seq_len=model_input.b_shared_seq_len if model_input.b_shared_seq_len is not None else dummy_1d, + out_b_shared_seq_len=out_b_shared_seq_len if out_b_shared_seq_len is not None else dummy_1d, + selected_mask=selected_mask_gpu, + selected_dst_pos=selected_dst_pos, + batch_size=old_batch_size, + HAS_INPUT_IDS=model_input.input_ids is not None, + HAS_B_POSITION_DELTA=model_input.b_position_delta is not None, + HAS_MEM_INDEXES=model_input.mem_indexes is not None, + HAS_B_SHARED_SEQ_LEN=model_input.b_shared_seq_len is not None, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=8, + num_stages=1, + ) + + model_input.input_ids = out_input_ids + model_input.b_req_idx = out_b_req_idx + model_input.b_mtp_index = out_b_mtp_index + model_input.b_seq_len = out_b_seq_len + model_input.b_position_delta = out_b_position_delta + model_input.mem_indexes = out_mem_indexes + model_input.b_shared_seq_len = out_b_shared_seq_len + model_input.b_mark_shared_group = _rebuild_trimmed_mtp_b_mark_shared_group_from_b_req_idx(out_b_req_idx) + + if model_input.input_ids is not None: + assert model_input.input_ids.shape[0] == dynamic_batch_size + if model_input.mtp_draft_input_hiddens is not None: + assert model_input.mtp_draft_input_hiddens.is_cuda + model_input.mtp_draft_input_hiddens = _pack_selected_rows_2d( + model_input.mtp_draft_input_hiddens, + selected_mask_gpu, + selected_dst_pos, + dynamic_batch_size, + ) + if model_input.b_position_delta is not None: + assert model_input.b_position_delta.shape[0] == dynamic_batch_size + if model_input.mem_indexes is not None: + assert model_input.mem_indexes.shape[0] == dynamic_batch_size + model_input.batch_size = dynamic_batch_size + + return model_input + + +def prepare_dynamic_mtp_model_input( + model_input: ModelInput, + req_num: int, + dynamic_batch_size: int, + req_to_next_token_ids: torch.Tensor, + req_to_next_token_probs: Optional[torch.Tensor] = None, + verify_step: Optional[int] = None, + use_prefix_selection: bool = False, +): + if req_to_next_token_probs is None: + selected_mask = torch.ones((model_input.batch_size,), dtype=torch.int32, device="cuda") + return model_input, selected_mask + + req_num = int(req_num) + dynamic_batch_size = int(dynamic_batch_size) + assert not model_input.is_prefill, "trim_dynamic_mtp_model_input only supports decode inputs" + assert dynamic_batch_size >= req_num + assert dynamic_batch_size <= model_input.batch_size + mtp_step = int(get_env_start_args().mtp_step) + verify_step = mtp_step if verify_step is None else int(verify_step) + assert 0 <= verify_step <= mtp_step + assert model_input.batch_size == req_num * (mtp_step + 1) + assert dynamic_batch_size <= req_num * (verify_step + 1) + + # ! 在一个CUDA流上面的GPU操作会自动串行化,因此不需要额外同步 + # ! model_input必须在GPU上,才能高效进行trim操作 + model_input.to_cuda() + + if use_prefix_selection: + assert dynamic_batch_size == req_num * (verify_step + 1) + selected_mask_gpu = (model_input.b_mtp_index <= verify_step).to(dtype=torch.int32) + else: + selected_mask_gpu = sample_dynamic_mtp_req_mask( + dynamic_batch_size=dynamic_batch_size, + b_req_idx=model_input.b_req_idx, + req_to_next_token_probs=req_to_next_token_probs, + mtp_step=mtp_step, + verify_step=verify_step, + ) + + model_input = _trim_decode_model_input_inplace( + model_input=model_input, + selected_mask_gpu=selected_mask_gpu, + dynamic_batch_size=dynamic_batch_size, + ) + # Keep CPU mem_indexes unfiltered here. Copying selected_mask_gpu back to + # CPU in this hot path synchronizes the overlap stream; the router frees + # unselected/rejected CPU mem indexes after its existing async mask copy is + # consumed. Decode only needs b_position_delta on device, so placeholder + # multimodal metadata keeps padded ModelInput checks shape-consistent. + if model_input.multimodal_params is not None: + # Read-only placeholders: avoid rebuilding hundreds of nested Python + # objects on every compacted decode iteration. + empty_multimodal_params = {"images": [], "audios": []} + model_input.multimodal_params = [empty_multimodal_params] * dynamic_batch_size + + # ! 现在这些值没有实际作用,同时修改这些值还会导致阻塞操作,因此暂时不修改这些值了 + # model_input.total_token_num = int(model_input.b_seq_len.sum().item()) + # model_input.max_kv_seq_len = int(model_input.b_seq_len.max().item()) + model_input.max_q_seq_len = 1 + return model_input, selected_mask_gpu + + @triton.jit def _fwd_kernel_gen_b_req_mtp_start_loc( b_mtp_index, b_req_mtp_start_loc, - num_reqs: tl.constexpr, - batch_size: tl.constexpr, + batch_size, BLOCK_SIZE: tl.constexpr, ): offset = tl.arange(0, BLOCK_SIZE) @@ -172,7 +573,6 @@ def gen_b_req_mtp_start_loc(b_mtp_index: torch.Tensor, num_reqs: int): _fwd_kernel_gen_b_req_mtp_start_loc[grid]( b_mtp_index=b_mtp_index, b_req_mtp_start_loc=b_req_mtp_start_loc, - num_reqs=num_reqs, batch_size=batch_size, BLOCK_SIZE=BLOCK_SIZE, num_warps=8, @@ -249,36 +649,3 @@ def linear_att_mtp_state_index_update( num_warps=num_warps, num_stages=1, ) - - -def test_mtp_verify(): - req_to_next_token_ids = torch.tensor( - [[1, 2, -2, -1, -1], [1, 2, 0, -1, -1], [1, 3, 4, 4, 5]], dtype=torch.int64, device="cuda" - ) - b_req_idx = torch.tensor([0, 0, 2, 2, 2], dtype=torch.int32, device="cuda") - b_req_mtp_start_loc = torch.tensor([0, 2], dtype=torch.int32, device="cuda") - new_next_token_ids = torch.tensor([1, 4, 2, 4, 13], dtype=torch.int64, device="cuda") - all_next_token_ids = torch.tensor( - [[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], [13, 14, 15]], dtype=torch.int64, device="cuda" - ) - mtp_accept_len, accepted_index = mtp_verify( - req_to_next_token_ids, b_req_mtp_start_loc, new_next_token_ids, b_req_idx - ) - mtp_scatter_next_token_ids( - req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, mtp_accept_len - ) - print(mtp_accept_len) - print(req_to_next_token_ids) - print(accepted_index) - - -def test_gen_b_req_mtp_start_loc(): - b_mtp_index = torch.tensor([0, 1, 0, 1, 2], dtype=torch.int32, device="cuda") - gt_output = torch.where(b_mtp_index == 0)[0] - b_req_mtp_start_loc = gen_b_req_mtp_start_loc(b_mtp_index, 2) - print(b_req_mtp_start_loc, gt_output) - - -if __name__ == "__main__": - test_mtp_verify() - # test_gen_b_req_mtp_start_loc() diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 070da7412f..f248cdbb39 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -7,7 +7,7 @@ from typing import List, Optional, TYPE_CHECKING from lightllm.common.basemodel.triton_kernel.gen_sampling_params import token_id_counter from lightllm.common.basemodel.triton_kernel.gen_sampling_params import update_req_to_token_id_counter -from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.envs_utils import get_env_start_args, enable_dynamic_mtp_verify from lightllm.utils.config_utils import get_vocab_size from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.common.linear_att_cache_manager.layer_cache import LayerCache @@ -116,11 +116,21 @@ def __init__(self, max_request_num): self.req_to_presence_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") self.req_to_frequency_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") self.req_to_repetition_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") + assert get_env_start_args().mtp_step <= 15, "mtp_step must be less than or equal to 15" self.req_to_next_token_ids = torch.zeros( - (max_request_num + 1, 8), + (max_request_num + 1, 16), dtype=torch.int64, device="cuda", ) + if enable_dynamic_mtp_verify(): + self.req_to_next_token_probs = torch.zeros( + (max_request_num + 1, 16), + dtype=torch.float32, + device="cuda", + ) + else: + self.req_to_next_token_probs = None + self.req_to_exponential_decay_length_penalty = torch.zeros( max_request_num + 1, dtype=torch.float32, device="cuda" ) @@ -137,6 +147,9 @@ def __init__(self, max_request_num): def init_req_sampling_params(self, req: "InferReq"): shm_param = req.sampling_param.shm_param self.req_to_next_token_ids[req.req_idx][0:1].fill_(req.get_last_gen_token()) + if enable_dynamic_mtp_verify(): + self.req_to_next_token_probs[req.req_idx].fill_(0.0) + self.req_to_next_token_probs[req.req_idx][0:1].fill_(1.0) self.req_to_presence_penalty[req.req_idx].fill_(shm_param.presence_penalty) self.req_to_frequency_penalty[req.req_idx].fill_(shm_param.frequency_penalty) self.req_to_repetition_penalty[req.req_idx].fill_(shm_param.repetition_penalty) diff --git a/lightllm/common/speculative/__init__.py b/lightllm/common/speculative/__init__.py new file mode 100644 index 0000000000..6507ee69ab --- /dev/null +++ b/lightllm/common/speculative/__init__.py @@ -0,0 +1,21 @@ +from .config import ( + SpeculativeConfig, + get_dspark_family_block_size, + is_dspark_draft_config, + is_eagle3_draft_config, + is_gemma4_dspark_draft_config, + is_qwen3_dflash_draft_config, + is_qwen3_dspark_draft_config, + validate_dspark_family_draft_config, +) + +__all__ = [ + "SpeculativeConfig", + "get_dspark_family_block_size", + "is_dspark_draft_config", + "is_eagle3_draft_config", + "is_gemma4_dspark_draft_config", + "is_qwen3_dflash_draft_config", + "is_qwen3_dspark_draft_config", + "validate_dspark_family_draft_config", +] diff --git a/lightllm/common/speculative/config.py b/lightllm/common/speculative/config.py new file mode 100644 index 0000000000..9c92dde73d --- /dev/null +++ b/lightllm/common/speculative/config.py @@ -0,0 +1,226 @@ +from dataclasses import dataclass +from typing import Any, Mapping, Optional + + +VANILLA_SPEC_MODES = frozenset({"vanilla_with_att", "vanilla_no_att", "qwen3next_vanilla"}) +EAGLE_SPEC_MODES = frozenset({"eagle_with_att", "eagle_no_att", "eagle3", "qwen3next_eagle"}) +BLOCK_SPEC_MODES = frozenset({"dspark", "dflash"}) +SPEC_MODES = VANILLA_SPEC_MODES | EAGLE_SPEC_MODES | BLOCK_SPEC_MODES + +ATTENTION_SPEC_MODES = frozenset({"vanilla_with_att", "eagle_with_att", "eagle3", "dspark", "dflash"}) +NO_ATTENTION_SPEC_MODES = frozenset({"vanilla_no_att", "eagle_no_att", "qwen3next_vanilla", "qwen3next_eagle"}) +TARGET_HIDDEN_SPEC_MODES = frozenset({"eagle3", "dspark", "dflash"}) +QWEN3_DFLASH_ARCHITECTURES = frozenset({"Qwen3DFlashModel", "Qwen3DSparkModel"}) +QWEN3_DSPARK_ARCHITECTURES = frozenset({"Qwen3DSparkModel"}) +GEMMA4_DSPARK_ARCHITECTURES = frozenset({"Gemma4DSparkModel"}) +DSPARK_FAMILY_ARCHITECTURES = QWEN3_DFLASH_ARCHITECTURES | GEMMA4_DSPARK_ARCHITECTURES +DSPARK_MARKOV_HEAD_TYPES = frozenset({"vanilla", "gated", "rnn"}) + + +@dataclass(frozen=True) +class SpeculativeConfig: + """Normalized view of speculative decoding mode flags.""" + + mode: Optional[str] + step: int + dynamic_verify: bool = False + + @classmethod + def from_args(cls, args: Any, dynamic_verify: Optional[bool] = None) -> "SpeculativeConfig": + mode = getattr(args, "mtp_mode", None) + if dynamic_verify is None: + dynamic_verify = bool(getattr(args, "mtp_dynamic_verify", False)) + if mode == "dspark": + dynamic_verify = True + elif mode == "dflash": + dynamic_verify = False + return cls( + mode=mode, + step=int(getattr(args, "mtp_step", 0)), + dynamic_verify=dynamic_verify, + ) + + @property + def enabled(self) -> bool: + return self.mode is not None + + @property + def is_vanilla(self) -> bool: + return self.mode in VANILLA_SPEC_MODES + + @property + def is_eagle(self) -> bool: + return self.mode in EAGLE_SPEC_MODES + + @property + def is_eagle3(self) -> bool: + return self.mode == "eagle3" + + @property + def is_dspark(self) -> bool: + return self.mode == "dspark" + + @property + def is_dflash(self) -> bool: + return self.mode == "dflash" + + @property + def uses_block_draft_model(self) -> bool: + return self.mode in BLOCK_SPEC_MODES + + @property + def needs_target_layer_hidden(self) -> bool: + return self.mode in TARGET_HIDDEN_SPEC_MODES + + @property + def uses_attention_draft(self) -> bool: + return self.mode in ATTENTION_SPEC_MODES + + @property + def uses_no_attention_draft(self) -> bool: + return self.mode in NO_ATTENTION_SPEC_MODES + + @property + def uses_chained_draft_models(self) -> bool: + return self.mode in VANILLA_SPEC_MODES + + @property + def uses_recurrent_draft_model(self) -> bool: + return self.mode in EAGLE_SPEC_MODES + + @property + def draft_model_count(self) -> int: + if not self.enabled: + return 0 + return 1 if (self.uses_recurrent_draft_model or self.uses_block_draft_model) else self.step + + @property + def needs_draft_vocab_mapping(self) -> bool: + return self.is_eagle3 + + def get_decode_graph_mtp_step(self, *, model_config: Mapping[str, Any], is_draft_model: bool) -> int: + if (self.is_dflash or self.is_dspark) and is_draft_model: + return int(model_config["block_size"]) - 1 + if self.is_eagle3 and self.dynamic_verify: + # Dynamic Eagle3 physically compacts target rows and recurrent + # draft rows, so graph shapes must be available at unit batch + # granularity instead of only at multiples of the exposed depth. + return 0 + return self.step + + def get_decode_graph_warmup_mtp_step(self, *, model_config: Mapping[str, Any], is_draft_model: bool) -> int: + if self.is_eagle3 and self.dynamic_verify: + # Preserve representative MTP indices/shared-group metadata while + # retaining unit-granularity graph shape capture. + return min(3, self.step) + return self.get_decode_graph_mtp_step( + model_config=model_config, + is_draft_model=is_draft_model, + ) + + def validate(self) -> None: + if not self.enabled: + assert self.step == 0 + return + + assert self.mode in SPEC_MODES, f"unsupported speculative mode {self.mode}" + if not self.uses_block_draft_model: + assert self.step > 0 + else: + assert self.step >= 0 + if self.is_dspark: + assert self.dynamic_verify, "DSpark mode requires dynamic verify scheduling" + if self.uses_chained_draft_models: + assert self.draft_model_count == self.step + else: + assert self.draft_model_count == 1 + + +def is_eagle3_draft_config(config: Mapping[str, Any]) -> bool: + architectures = config.get("architectures", []) + return config.get("model_type") == "llama" or any( + architecture in ["Eagle3Speculator", "Qwen3Eagle3Model"] for architecture in architectures + ) + + +def is_dspark_draft_config(config: Mapping[str, Any]) -> bool: + architectures = config.get("architectures", []) + return any(architecture in DSPARK_FAMILY_ARCHITECTURES for architecture in architectures) + + +def is_qwen3_dflash_draft_config(config: Mapping[str, Any]) -> bool: + architectures = config.get("architectures", []) + return any(architecture in QWEN3_DFLASH_ARCHITECTURES for architecture in architectures) + + +def is_qwen3_dspark_draft_config(config: Mapping[str, Any]) -> bool: + architectures = config.get("architectures", []) + return any(architecture in QWEN3_DSPARK_ARCHITECTURES for architecture in architectures) + + +def is_gemma4_dspark_draft_config(config: Mapping[str, Any]) -> bool: + architectures = config.get("architectures", []) + return any(architecture in GEMMA4_DSPARK_ARCHITECTURES for architecture in architectures) + + +def validate_dspark_family_draft_config( + config: Mapping[str, Any], + *, + require_confidence_head: bool = False, +) -> None: + """Validate DFlash/DSpark checkpoint fields consumed by LightLLM serving.""" + + assert is_dspark_draft_config(config), f"unsupported DFlash/DSpark architecture: {config.get('architectures')}" + + block_size = int(config.get("block_size", 0)) + assert block_size > 0, "DFlash/DSpark draft config must provide positive block_size" + + target_layer_ids = config.get("target_layer_ids") + assert ( + isinstance(target_layer_ids, (list, tuple)) and len(target_layer_ids) > 0 + ), "DFlash/DSpark draft config must provide non-empty target_layer_ids" + previous_layer_id = None + for raw_layer_id in target_layer_ids: + layer_id = int(raw_layer_id) + assert layer_id >= 0, ( + "LightLLM DFlash/DSpark serving expects decoder-layer target_layer_ids; " + "embedding-output layer_id=-1 is not supported" + ) + assert ( + previous_layer_id is None or layer_id > previous_layer_id + ), "DFlash/DSpark target_layer_ids must be strictly increasing" + previous_layer_id = layer_id + + assert "mask_token_id" in config, "DFlash/DSpark draft config must provide mask_token_id" + assert int(config["mask_token_id"]) >= 0, "DFlash/DSpark mask_token_id must be non-negative" + + markov_rank = int(config.get("markov_rank", 0)) + assert markov_rank >= 0, f"DFlash/DSpark markov_rank must be >= 0, got {markov_rank}" + if markov_rank > 0: + markov_head_type = str(config.get("markov_head_type", "")).lower() + assert ( + markov_head_type in DSPARK_MARKOV_HEAD_TYPES + ), f"unsupported DFlash/DSpark markov_head_type {markov_head_type!r}" + + enable_confidence_head = bool(config.get("enable_confidence_head", False)) + if require_confidence_head: + assert enable_confidence_head, "DSpark dynamic scheduling requires enable_confidence_head=true" + if enable_confidence_head: + assert ( + "confidence_head_with_markov" in config + ), "confidence_head_with_markov must be provided when enable_confidence_head is true" + if bool(config.get("confidence_head_with_markov", False)): + assert markov_rank > 0, "confidence_head_with_markov requires markov_rank > 0" + return + + +def get_dspark_family_block_size( + config: Mapping[str, Any], + *, + require_confidence_head: bool = False, +) -> int: + validate_dspark_family_draft_config( + config, + require_confidence_head=require_confidence_head, + ) + return int(config["block_size"]) diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H800/_fwd_kernel_mtp_diverse_stage1_single_token:v1/{block_seq=256,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H800/_fwd_kernel_mtp_diverse_stage1_single_token:v1/{block_seq=256,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H800.json new file mode 100644 index 0000000000..6e0ce74445 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H800/_fwd_kernel_mtp_diverse_stage1_single_token:v1/{block_seq=256,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H800.json @@ -0,0 +1,326 @@ +{ + "1000000032": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 4 + }, + "1000000064": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 8 + }, + "1000000128": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 8 + }, + "1000000256": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 8 + }, + "1000000512": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 8 + }, + "1000001024": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 8 + }, + "1000002048": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 8 + }, + "1000008192": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 8 + }, + "1000016384": { + "BLOCK_BATCH": 4, + "BLOCK_N": 32, + "num_stages": 3, + "num_warps": 4 + }, + "128000000032": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 2 + }, + "128000000064": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2 + }, + "128000000128": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 2 + }, + "128000000256": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 4, + "num_warps": 2 + }, + "128000000512": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 3, + "num_warps": 2 + }, + "128000001024": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 4 + }, + "128000002048": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 4 + }, + "128000008192": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 2 + }, + "128000016384": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 2 + }, + "16000000032": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 4 + }, + "16000000064": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 8 + }, + "16000000128": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 8 + }, + "16000000256": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 8 + }, + "16000000512": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 4 + }, + "16000001024": { + "BLOCK_BATCH": 4, + "BLOCK_N": 32, + "num_stages": 4, + "num_warps": 4 + }, + "16000002048": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 4, + "num_warps": 2 + }, + "16000008192": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 2 + }, + "16000016384": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 4 + }, + "32000000032": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 4 + }, + "32000000064": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 4 + }, + "32000000128": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 4 + }, + "32000000256": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 4 + }, + "32000000512": { + "BLOCK_BATCH": 4, + "BLOCK_N": 32, + "num_stages": 4, + "num_warps": 4 + }, + "32000001024": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 3, + "num_warps": 2 + }, + "32000002048": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 3, + "num_warps": 2 + }, + "32000008192": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 2 + }, + "32000016384": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 2 + }, + "64000000032": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 2 + }, + "64000000064": { + "BLOCK_BATCH": 4, + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 4 + }, + "64000000128": { + "BLOCK_BATCH": 4, + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 4 + }, + "64000000256": { + "BLOCK_BATCH": 4, + "BLOCK_N": 32, + "num_stages": 4, + "num_warps": 2 + }, + "64000000512": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 3, + "num_warps": 2 + }, + "64000001024": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 3, + "num_warps": 2 + }, + "64000002048": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 2 + }, + "64000008192": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 2 + }, + "64000016384": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 2 + }, + "8000000032": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 8 + }, + "8000000064": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 4 + }, + "8000000128": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 8 + }, + "8000000256": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 8 + }, + "8000000512": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 8 + }, + "8000001024": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 4, + "num_warps": 4 + }, + "8000002048": { + "BLOCK_BATCH": 4, + "BLOCK_N": 32, + "num_stages": 3, + "num_warps": 4 + }, + "8000008192": { + "BLOCK_BATCH": 4, + "BLOCK_N": 16, + "num_stages": 4, + "num_warps": 2 + }, + "8000016384": { + "BLOCK_BATCH": 4, + "BLOCK_N": 64, + "num_stages": 3, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..b730bccb8d --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,254 @@ +{ + "1000000032": { + "BLOCK_N": 16, + "num_stages": 3, + "num_warps": 2, + "warp_specialize": true + }, + "1000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 4, + "warp_specialize": false + }, + "1000000128": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 8, + "warp_specialize": true + }, + "1000000256": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "1000000512": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 8, + "warp_specialize": false + }, + "1000001024": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "1000002048": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "128000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "128000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "128000000128": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "128000000256": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "128000000512": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "128000001024": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "128000002048": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "16000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "16000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "16000000128": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "16000000256": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "16000000512": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "16000001024": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "16000002048": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "32000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "32000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "32000000128": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "32000000256": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "32000000512": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "32000001024": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "32000002048": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "64000000128": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "64000000256": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000000512": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000001024": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000002048": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "8000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "8000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "8000000128": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 4, + "warp_specialize": false + }, + "8000000256": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "8000000512": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 8, + "warp_specialize": true + }, + "8000001024": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "8000002048": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=8,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=8,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..d04a6d8fe2 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=8,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,254 @@ +{ + "1000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "1000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 4, + "warp_specialize": false + }, + "1000000128": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "1000000256": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "1000000512": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 4, + "warp_specialize": false + }, + "1000001024": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "1000002048": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "128000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "128000000064": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "128000000128": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "128000000256": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "128000000512": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "128000001024": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "128000002048": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "16000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "16000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "16000000128": { + "BLOCK_N": 16, + "num_stages": 3, + "num_warps": 2, + "warp_specialize": false + }, + "16000000256": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "16000000512": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "16000001024": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "16000002048": { + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 4, + "warp_specialize": false + }, + "32000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "32000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "32000000128": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "32000000256": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "32000000512": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "32000001024": { + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 4, + "warp_specialize": false + }, + "32000002048": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000000064": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "64000000128": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "64000000256": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000000512": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "64000001024": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "64000002048": { + "BLOCK_N": 64, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "8000000032": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "8000000064": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "8000000128": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "8000000256": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": true + }, + "8000000512": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "8000001024": { + "BLOCK_N": 16, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + }, + "8000002048": { + "BLOCK_N": 32, + "num_stages": 2, + "num_warps": 2, + "warp_specialize": false + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..fcb4db67b1 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,26 @@ +{ + "1000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "128000000032": { + "num_stages": 1, + "num_warps": 2 + }, + "16000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "32000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "64000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "8000000128": { + "num_stages": 1, + "num_warps": 2 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..d5a8839f80 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,26 @@ +{ + "1000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "128000000032": { + "num_stages": 1, + "num_warps": 4 + }, + "16000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "32000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "64000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "8000000128": { + "num_stages": 1, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..d75cfb554e --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,26 @@ +{ + "1000000128": { + "num_stages": 1, + "num_warps": 2 + }, + "128000000032": { + "num_stages": 1, + "num_warps": 2 + }, + "16000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "32000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "64000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "8000000128": { + "num_stages": 1, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..a4b1a211cd --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,26 @@ +{ + "1000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "128000000032": { + "num_stages": 1, + "num_warps": 4 + }, + "16000000128": { + "num_stages": 1, + "num_warps": 2 + }, + "32000000064": { + "num_stages": 1, + "num_warps": 2 + }, + "64000000064": { + "num_stages": 1, + "num_warps": 2 + }, + "8000000128": { + "num_stages": 1, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..d5a8839f80 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,26 @@ +{ + "1000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "128000000032": { + "num_stages": 1, + "num_warps": 4 + }, + "16000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "32000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "64000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "8000000128": { + "num_stages": 1, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..42975b4b9f --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,26 @@ +{ + "1000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "128000000032": { + "num_stages": 1, + "num_warps": 4 + }, + "16000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "32000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "64000000064": { + "num_stages": 1, + "num_warps": 2 + }, + "8000000128": { + "num_stages": 1, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..0220ae94ed --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,26 @@ +{ + "1000000128": { + "num_stages": 1, + "num_warps": 2 + }, + "128000000032": { + "num_stages": 1, + "num_warps": 2 + }, + "16000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "32000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "64000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "8000000128": { + "num_stages": 1, + "num_warps": 2 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json new file mode 100644 index 0000000000..13615a6096 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json @@ -0,0 +1,26 @@ +{ + "1000000128": { + "num_stages": 1, + "num_warps": 4 + }, + "128000000032": { + "num_stages": 1, + "num_warps": 2 + }, + "16000000128": { + "num_stages": 1, + "num_warps": 2 + }, + "32000000064": { + "num_stages": 1, + "num_warps": 2 + }, + "64000000064": { + "num_stages": 1, + "num_warps": 4 + }, + "8000000128": { + "num_stages": 1, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/models/__init__.py b/lightllm/models/__init__.py index f619b1d88f..d56b17608a 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -1,46 +1,173 @@ -from lightllm.models.mixtral.model import MixtralTpPartModel -from lightllm.models.bloom.model import BloomTpPartModel -from lightllm.models.llama.model import LlamaTpPartModel -from lightllm.models.starcoder.model import StarcoderTpPartModel -from lightllm.models.starcoder2.model import Starcoder2TpPartModel -from lightllm.models.qwen.model import QWenTpPartModel -from lightllm.models.qwen2.model import Qwen2TpPartModel -from lightllm.models.qwen3.model import Qwen3TpPartModel -from lightllm.models.qwen3_moe.model import Qwen3MOEModel -from lightllm.models.qwen3next.model import Qwen3NextTpPartModel -from lightllm.models.internlm.model import InternlmTpPartModel -from lightllm.models.stablelm.model import StablelmTpPartModel -from lightllm.models.internlm2.model import Internlm2TpPartModel -from lightllm.models.internlm2_reward.model import Internlm2RewardTpPartModel -from lightllm.models.mistral.model import MistralTpPartModel -from lightllm.models.minicpm.model import MiniCPMTpPartModel -from lightllm.models.llava.model import LlavaTpPartModel -from lightllm.models.qwen_vl.model import QWenVLTpPartModel -from lightllm.models.gemma_2b.model import Gemma_2bTpPartModel -from lightllm.models.phi3.model import Phi3TpPartModel -from lightllm.models.deepseek2.model import Deepseek2TpPartModel -from lightllm.models.deepseek3_2.model import Deepseek3_2TpPartModel -from lightllm.models.glm4_moe_lite.model import Glm4MoeLiteTpPartModel -from lightllm.models.internvl.model import ( - InternVLLlamaTpPartModel, - InternVLPhi3TpPartModel, - InternVLQwen2TpPartModel, - InternVLDeepSeek2TpPartModel, -) -from lightllm.models.internvl.model import InternVLInternlm2TpPartModel -from lightllm.models.qwen2_vl.model import Qwen2VLTpPartModel -from lightllm.models.qwen2_reward.model import Qwen2RewardTpPartModel -from lightllm.models.qwen3_vl.model import Qwen3VLTpPartModel -from lightllm.models.qwen3_vl_moe.model import Qwen3VLMOETpPartModel -from lightllm.models.gemma3.model import Gemma3TpPartModel -from lightllm.models.gemma4.model import Gemma4TpPartModel -from lightllm.models.tarsier2.model import ( - Tarsier2Qwen2TpPartModel, - Tarsier2Qwen2VLTpPartModel, - Tarsier2LlamaTpPartModel, -) -from lightllm.models.gpt_oss.model import GptOssTpPartModel -from lightllm.models.qwen3_omni_moe_thinker.model import Qwen3OmniMOETpPartModel -from lightllm.models.qwen3_5.model import Qwen3_5TpPartModel -from lightllm.models.qwen3_5_moe.model import Qwen3_5MOETpPartModel -from .registry import get_model, get_model_class +from importlib import import_module + +from .registry import get_model as _registry_get_model +from .registry import get_model_class as _registry_get_model_class + + +_MODEL_EXPORTS = { + "MixtralTpPartModel": ("lightllm.models.mixtral.model", "MixtralTpPartModel"), + "BloomTpPartModel": ("lightllm.models.bloom.model", "BloomTpPartModel"), + "LlamaTpPartModel": ("lightllm.models.llama.model", "LlamaTpPartModel"), + "StarcoderTpPartModel": ("lightllm.models.starcoder.model", "StarcoderTpPartModel"), + "Starcoder2TpPartModel": ("lightllm.models.starcoder2.model", "Starcoder2TpPartModel"), + "QWenTpPartModel": ("lightllm.models.qwen.model", "QWenTpPartModel"), + "Qwen2TpPartModel": ("lightllm.models.qwen2.model", "Qwen2TpPartModel"), + "Qwen3TpPartModel": ("lightllm.models.qwen3.model", "Qwen3TpPartModel"), + "Qwen3MOEModel": ("lightllm.models.qwen3_moe.model", "Qwen3MOEModel"), + "Qwen3NextTpPartModel": ("lightllm.models.qwen3next.model", "Qwen3NextTpPartModel"), + "InternlmTpPartModel": ("lightllm.models.internlm.model", "InternlmTpPartModel"), + "StablelmTpPartModel": ("lightllm.models.stablelm.model", "StablelmTpPartModel"), + "Internlm2TpPartModel": ("lightllm.models.internlm2.model", "Internlm2TpPartModel"), + "Internlm2RewardTpPartModel": ( + "lightllm.models.internlm2_reward.model", + "Internlm2RewardTpPartModel", + ), + "MistralTpPartModel": ("lightllm.models.mistral.model", "MistralTpPartModel"), + "MiniCPMTpPartModel": ("lightllm.models.minicpm.model", "MiniCPMTpPartModel"), + "LlavaTpPartModel": ("lightllm.models.llava.model", "LlavaTpPartModel"), + "QWenVLTpPartModel": ("lightllm.models.qwen_vl.model", "QWenVLTpPartModel"), + "Gemma_2bTpPartModel": ("lightllm.models.gemma_2b.model", "Gemma_2bTpPartModel"), + "Phi3TpPartModel": ("lightllm.models.phi3.model", "Phi3TpPartModel"), + "Deepseek2TpPartModel": ("lightllm.models.deepseek2.model", "Deepseek2TpPartModel"), + "Deepseek3_2TpPartModel": ("lightllm.models.deepseek3_2.model", "Deepseek3_2TpPartModel"), + "Glm4MoeLiteTpPartModel": ( + "lightllm.models.glm4_moe_lite.model", + "Glm4MoeLiteTpPartModel", + ), + "InternVLLlamaTpPartModel": ("lightllm.models.internvl.model", "InternVLLlamaTpPartModel"), + "InternVLPhi3TpPartModel": ("lightllm.models.internvl.model", "InternVLPhi3TpPartModel"), + "InternVLQwen2TpPartModel": ("lightllm.models.internvl.model", "InternVLQwen2TpPartModel"), + "InternVLDeepSeek2TpPartModel": ( + "lightllm.models.internvl.model", + "InternVLDeepSeek2TpPartModel", + ), + "InternVLInternlm2TpPartModel": ( + "lightllm.models.internvl.model", + "InternVLInternlm2TpPartModel", + ), + "Qwen2VLTpPartModel": ("lightllm.models.qwen2_vl.model", "Qwen2VLTpPartModel"), + "Qwen2RewardTpPartModel": ("lightllm.models.qwen2_reward.model", "Qwen2RewardTpPartModel"), + "Qwen3VLTpPartModel": ("lightllm.models.qwen3_vl.model", "Qwen3VLTpPartModel"), + "Qwen3VLMOETpPartModel": ("lightllm.models.qwen3_vl_moe.model", "Qwen3VLMOETpPartModel"), + "Gemma3TpPartModel": ("lightllm.models.gemma3.model", "Gemma3TpPartModel"), + "Gemma4TpPartModel": ("lightllm.models.gemma4.model", "Gemma4TpPartModel"), + "Tarsier2Qwen2TpPartModel": ("lightllm.models.tarsier2.model", "Tarsier2Qwen2TpPartModel"), + "Tarsier2Qwen2VLTpPartModel": ( + "lightllm.models.tarsier2.model", + "Tarsier2Qwen2VLTpPartModel", + ), + "Tarsier2LlamaTpPartModel": ("lightllm.models.tarsier2.model", "Tarsier2LlamaTpPartModel"), + "GptOssTpPartModel": ("lightllm.models.gpt_oss.model", "GptOssTpPartModel"), + "Qwen3OmniMOETpPartModel": ( + "lightllm.models.qwen3_omni_moe_thinker.model", + "Qwen3OmniMOETpPartModel", + ), + "Qwen3_5TpPartModel": ("lightllm.models.qwen3_5.model", "Qwen3_5TpPartModel"), + "Qwen3_5MOETpPartModel": ("lightllm.models.qwen3_5_moe.model", "Qwen3_5MOETpPartModel"), +} + +_MODEL_TYPE_REGISTRY_MODULES = { + "starcoder2": ("lightllm.models.starcoder2.model",), + "internlm2": ("lightllm.models.internlm2.model",), + "llava": ("lightllm.models.llava.model",), + "qwen": ("lightllm.models.qwen.model",), + "qwen2": ("lightllm.models.qwen2.model",), + "qwen2_vl": ("lightllm.models.qwen2_vl.model",), + "qwen2_5_vl": ("lightllm.models.qwen2_vl.model",), + "qwen3": ("lightllm.models.qwen3.model",), + "qwen3_moe": ("lightllm.models.qwen3_moe.model",), + "qwen3_next": ("lightllm.models.qwen3next.model",), + "qwen3_vl": ("lightllm.models.qwen3_vl.model",), + "qwen3_vl_moe": ("lightllm.models.qwen3_vl_moe.model",), + "qwen3_omni_moe": ("lightllm.models.qwen3_omni_moe_thinker.model",), + "qwen3_5": ("lightllm.models.qwen3_5.model",), + "qwen3_5_moe": ("lightllm.models.qwen3_5_moe.model",), + "deepseek_v2": ("lightllm.models.deepseek2.model",), + "deepseek_v3": ("lightllm.models.deepseek2.model",), + "deepseek_v32": ("lightllm.models.deepseek3_2.model",), + "glm4_moe_lite": ("lightllm.models.glm4_moe_lite.model",), + "bloom": ("lightllm.models.bloom.model",), + "gpt_bigcode": ("lightllm.models.starcoder.model",), + "minicpm": ("lightllm.models.minicpm.model",), + "gemma3": ("lightllm.models.gemma3.model",), + "gemma": ("lightllm.models.gemma_2b.model",), + "gemma4": ("lightllm.models.gemma4.model",), + "internlm": ("lightllm.models.internlm.model",), + "stablelm": ("lightllm.models.stablelm.model",), + "mistral": ("lightllm.models.mistral.model",), + "gpt_oss": ("lightllm.models.gpt_oss.model",), + "phi3": ("lightllm.models.phi3.model",), + "llama": ("lightllm.models.llama.model",), + "mixtral": ("lightllm.models.mixtral.model",), + "internvl_chat": ("lightllm.models.internvl.model",), +} +_bootstrapped_registry_modules = set() + + +def _load_model_attr(name): + module_name, attr_name = _MODEL_EXPORTS[name] + value = getattr(import_module(module_name), attr_name) + globals()[name] = value + return value + + +def _has_architecture(model_cfg: dict, name: str) -> bool: + return any(name in architecture for architecture in model_cfg.get("architectures", [])) + + +def _llava_text_model_type(model_cfg: dict) -> str: + return model_cfg.get("llm_config", {}).get("model_type", "") or model_cfg.get("text_config", {}).get( + "model_type", "" + ) + + +def _registry_modules_for_model_cfg(model_cfg: dict): + model_type = str(model_cfg.get("model_type", "")) + + module_names = _MODEL_TYPE_REGISTRY_MODULES.get(model_type) + if module_names is None: + # Leave already-registered plugin/custom models available, but avoid + # importing every built-in module just to produce an unsupported-model + # error. Some built-ins have optional multimodal dependencies. + return () + + if model_type == "qwen" and "visual" in model_cfg: + module_names = module_names + ("lightllm.models.qwen_vl.model",) + elif model_type == "qwen2" and _has_architecture(model_cfg, "RewardModel"): + module_names = module_names + ("lightllm.models.qwen2_reward.model",) + elif model_type == "internlm2" and _has_architecture(model_cfg, "RewardModel"): + module_names = module_names + ("lightllm.models.internlm2_reward.model",) + elif model_type == "llava" and _llava_text_model_type(model_cfg) in {"qwen2", "qwen2_vl", "llama"}: + module_names = module_names + ("lightllm.models.tarsier2.model",) + + return module_names + + +def _ensure_model_registry_bootstrapped(model_cfg: dict) -> None: + module_names = _registry_modules_for_model_cfg(model_cfg) + + for module_name in module_names: + if module_name in _bootstrapped_registry_modules: + continue + import_module(module_name) + _bootstrapped_registry_modules.add(module_name) + return + + +def get_model(model_cfg: dict, model_kvargs: dict): + _ensure_model_registry_bootstrapped(model_cfg) + return _registry_get_model(model_cfg, model_kvargs) + + +def get_model_class(model_cfg: dict): + _ensure_model_registry_bootstrapped(model_cfg) + return _registry_get_model_class(model_cfg) + + +def __getattr__(name): + if name in _MODEL_EXPORTS: + return _load_model_attr(name) + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + +__all__ = ["get_model", "get_model_class"] + list(_MODEL_EXPORTS) diff --git a/lightllm/models/deepseek_mtp/model.py b/lightllm/models/deepseek_mtp/model.py index e2b2a56137..c6ca8dad53 100644 --- a/lightllm/models/deepseek_mtp/model.py +++ b/lightllm/models/deepseek_mtp/model.py @@ -6,10 +6,6 @@ class Deepseek3MTPModel(Deepseek2TpPartModel): - - # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). - is_mtp_draft_model = True - pre_and_post_weight_class = Deepseek3MTPPreAndPostLayerWeight pre_layer_infer_class = Deepseek3MTPPreLayerInfer @@ -23,6 +19,9 @@ def _pre_init(self, kvargs: dict): self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") return + def _gen_special_model_input(self, token_num: int): + return self._gen_mtp_draft_special_model_input(token_num) + def _init_custom(self): self._cos_cached = self.main_model._cos_cached self._sin_cached = self.main_model._sin_cached diff --git a/lightllm/models/glm4_moe_lite_mtp/model.py b/lightllm/models/glm4_moe_lite_mtp/model.py index 2e4ba5c86b..95ebf23efd 100644 --- a/lightllm/models/glm4_moe_lite_mtp/model.py +++ b/lightllm/models/glm4_moe_lite_mtp/model.py @@ -9,10 +9,6 @@ class Glm4MoeLiteMTPModel(Glm4MoeLiteTpPartModel): - - # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). - is_mtp_draft_model = True - pre_and_post_weight_class = Glm4MoeLiteMTPPreAndPostLayerWeight pre_layer_infer_class = Deepseek3MTPPreLayerInfer @@ -24,6 +20,9 @@ def _pre_init(self, kvargs: dict): self.main_model: TpPartBaseModel = kvargs.pop("main_model") self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") + def _gen_special_model_input(self, token_num: int): + return self._gen_mtp_draft_special_model_input(token_num) + def _init_custom(self): self._cos_cached = self.main_model._cos_cached self._sin_cached = self.main_model._sin_cached diff --git a/lightllm/models/mistral_mtp/model.py b/lightllm/models/mistral_mtp/model.py index f17bc0a383..7c1ac50ba7 100644 --- a/lightllm/models/mistral_mtp/model.py +++ b/lightllm/models/mistral_mtp/model.py @@ -9,10 +9,6 @@ class MistralMTPModel(MistralTpPartModel): - - # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). - is_mtp_draft_model = True - pre_and_post_weight_class = MistralMTPPreAndPostLayerWeight pre_layer_infer_class = MistralMTPPreLayerInfer @@ -31,6 +27,9 @@ def _pre_init(self, kvargs: dict): self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") return + def _gen_special_model_input(self, token_num: int): + return self._gen_mtp_draft_special_model_input(token_num) + def _init_some_value(self): super()._init_some_value() self.layers_num = 1 diff --git a/lightllm/models/qwen3_dflash/__init__.py b/lightllm/models/qwen3_dflash/__init__.py new file mode 100644 index 0000000000..0f3fd678aa --- /dev/null +++ b/lightllm/models/qwen3_dflash/__init__.py @@ -0,0 +1,3 @@ +from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel + +__all__ = ["Qwen3DFlashModel"] diff --git a/lightllm/models/qwen3_dflash/infer_struct.py b/lightllm/models/qwen3_dflash/infer_struct.py new file mode 100644 index 0000000000..6ea456fddf --- /dev/null +++ b/lightllm/models/qwen3_dflash/infer_struct.py @@ -0,0 +1,40 @@ +import torch + +from lightllm.models.llama.infer_struct import LlamaInferStateInfo + + +class Qwen3DFlashInferStateInfo(LlamaInferStateInfo): + """DFlash metadata on top of the normal prefill state.""" + + def __init__(self): + super().__init__() + self.prefill_causal: bool = True + self.decode_causal: bool = True + self.decode_mtp_step: int = None + + def init_some_extra_state(self, model): + super().init_some_extra_state(model) + if self.is_prefill: + self.prefill_causal = False + else: + self.decode_causal = False + self.decode_mtp_step = model.block_size - 1 + self.is_draft_model = True + return + + @staticmethod + def build_draft_query_position_ids( + *, + selected_seq_len: torch.Tensor, + b_position_delta: torch.Tensor = None, + draft_step: int, + ) -> torch.Tensor: + offsets = torch.arange( + int(draft_step), + dtype=torch.long, + device=selected_seq_len.device, + ) + position_ids = selected_seq_len.to(torch.long).view(-1, 1) + offsets.view(1, -1) + if b_position_delta is not None: + position_ids = position_ids + b_position_delta.to(torch.long).view(-1, 1) + return position_ids diff --git a/lightllm/models/qwen3_dflash/layer_infer/__init__.py b/lightllm/models/qwen3_dflash/layer_infer/__init__.py new file mode 100644 index 0000000000..16e5628d81 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_infer/__init__.py @@ -0,0 +1,9 @@ +from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer +from lightllm.models.qwen3_dflash.layer_infer.pre_layer_infer import Qwen3DFlashPreLayerInfer +from lightllm.models.qwen3_dflash.layer_infer.transformer_layer_infer import Qwen3DFlashTransformerLayerInfer + +__all__ = [ + "Qwen3DFlashPostLayerInfer", + "Qwen3DFlashPreLayerInfer", + "Qwen3DFlashTransformerLayerInfer", +] diff --git a/lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py new file mode 100644 index 0000000000..714513e211 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py @@ -0,0 +1,28 @@ +import torch + +from lightllm.models.llama.layer_infer.post_layer_infer import LlamaPostLayerInfer + + +class Qwen3DFlashPostLayerInfer(LlamaPostLayerInfer): + def _is_commit_prefill(self, infer_state): + return infer_state.is_prefill and infer_state.mtp_draft_input_hiddens is not None + + def _tpsp_allgather(self, input: torch.Tensor, infer_state): + if self._is_commit_prefill(infer_state): + return input + return super()._tpsp_allgather(input=input, infer_state=infer_state) + + def token_forward(self, input_embdings: torch.Tensor, infer_state, layer_weight): + if self._is_commit_prefill(infer_state): + # Commit prefill only materializes draft KV. There is no LM head + # work to do, but BaseModel still expects a logits-shaped tensor. + return torch.empty( + (infer_state.input_ids.shape[0], 0), + dtype=input_embdings.dtype, + device=input_embdings.device, + ) + return super().token_forward( + input_embdings=input_embdings, + infer_state=infer_state, + layer_weight=layer_weight, + ) diff --git a/lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py b/lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py new file mode 100644 index 0000000000..b7e59529e7 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py @@ -0,0 +1,67 @@ +import torch + +from lightllm.models.llama.layer_infer.pre_layer_infer import LlamaPreLayerInfer +from lightllm.models.qwen3_dflash.layer_weights.pre_and_post_layer_weight import Qwen3DFlashPreAndPostLayerWeight + + +class Qwen3DFlashPreLayerInfer(LlamaPreLayerInfer): + """DFlash target-hidden projection plus normal token embedding.""" + + def __init__(self, network_config): + super().__init__(network_config) + self.eps_ = network_config["rms_norm_eps"] + self.hidden_size_ = network_config["hidden_size"] + return + + def project_target_hidden( + self, + *, + target_hidden_states: torch.Tensor, + layer_weight: Qwen3DFlashPreAndPostLayerWeight, + ) -> torch.Tensor: + if target_hidden_states.dim() == 2: + batch_size = target_hidden_states.shape[0] + context_len = 1 + flat_hidden = target_hidden_states + else: + assert target_hidden_states.dim() == 3 + batch_size, context_len, _ = target_hidden_states.shape + flat_hidden = target_hidden_states.reshape(batch_size * context_len, -1) + + projected = layer_weight.fc_weight_.mm(flat_hidden, use_custom_tensor_mananger=False) + projected = layer_weight.hidden_norm_weight_( + input=projected, + eps=self.eps_, + alloc_func=self.alloc_tensor, + ) + return projected.view(batch_size, context_len, self.hidden_size_) + + def context_forward( + self, + input_ids, + infer_state, + layer_weight: Qwen3DFlashPreAndPostLayerWeight, + ): + if infer_state.mtp_draft_input_hiddens is None: + return super().context_forward(input_ids, infer_state, layer_weight) + + return self.project_target_hidden( + target_hidden_states=infer_state.mtp_draft_input_hiddens, + layer_weight=layer_weight, + ).reshape(-1, self.hidden_size_) + + def token_forward( + self, + input_ids, + infer_state, + layer_weight: Qwen3DFlashPreAndPostLayerWeight, + ): + return super().token_forward(input_ids, infer_state, layer_weight) + + def decode_forward( + self, + input_ids, + infer_state, + layer_weight: Qwen3DFlashPreAndPostLayerWeight, + ): + return self.token_forward(input_ids, infer_state, layer_weight) diff --git a/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py new file mode 100644 index 0000000000..54680af804 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py @@ -0,0 +1,113 @@ +import torch + +from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer +from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd +from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo +from lightllm.models.qwen3_dflash.layer_weights.transformer_layer_weight import Qwen3DFlashTransformerLayerWeight + + +class Qwen3DFlashTransformerLayerInfer(LlamaTransformerLayerInfer): + """DFlash layer inference. + + The model path is built from two explicit layer primitives: + - commit accepted target hidden rows into draft KV + - run one non-causal draft block over prefix KV + scratch KV + """ + + def __init__(self, layer_num, network_config): + super().__init__(layer_num, network_config) + self.head_dim_ = network_config["head_dim"] + self.block_size_ = int(network_config["block_size"]) + return + + def context_forward( + self, + input_embdings: torch.Tensor, + infer_state: Qwen3DFlashInferStateInfo, + layer_weight: Qwen3DFlashTransformerLayerWeight, + ) -> torch.Tensor: + token_num, _ = input_embdings.shape + kv = layer_weight.kv_proj.mm(input_embdings, use_custom_tensor_mananger=False) + kv = kv.view(token_num, self.tp_k_head_num_ + self.tp_v_head_num_, self.head_dim_) + k = kv[:, : self.tp_k_head_num_, :] + v = kv[:, self.tp_k_head_num_ :, :] + k = layer_weight.k_norm_weight_( + input=k.reshape(-1, self.head_dim_), + eps=self.eps_, + alloc_func=torch.empty, + ).view(token_num, self.tp_k_head_num_, self.head_dim_) + rotary_emb_fwd( + k, + None, + infer_state.position_cos, + infer_state.position_sin, + ) + cache_kv = torch.cat([k, v], dim=1) + self._post_cache_kv(cache_kv.contiguous(), infer_state, layer_weight) + return input_embdings + + def token_forward( + self, + input_embdings: torch.Tensor, + infer_state: Qwen3DFlashInferStateInfo, + layer_weight: Qwen3DFlashTransformerLayerWeight, + ) -> torch.Tensor: + hidden_states = input_embdings.view(-1, self.block_size_, self.embed_dim_) + residual = hidden_states + q, cache_kv = self._get_qkv(hidden_states, infer_state, layer_weight) + batch_size, block_size, hidden_size = hidden_states.shape + self._post_cache_kv(cache_kv.contiguous(), infer_state, layer_weight) + o = self._token_attention_kernel(q, infer_state, layer_weight) + o = self._get_o(o, infer_state=infer_state, layer_weight=layer_weight) + hidden_states = residual + o.view(-1, block_size, hidden_size) + + residual = hidden_states + ffn_input = self._ffn_norm( + hidden_states.reshape(-1, hidden_size), + infer_state=infer_state, + layer_weight=layer_weight, + ) + ffn_out = self._ffn(ffn_input, infer_state=infer_state, layer_weight=layer_weight) + hidden_states = residual + ffn_out.view(-1, block_size, hidden_size) + return hidden_states.view(-1, self.embed_dim_) + + def _get_qkv(self, input, infer_state: Qwen3DFlashInferStateInfo, layer_weight: Qwen3DFlashTransformerLayerWeight): + hidden_states = input.view(-1, self.block_size_, self.embed_dim_) + batch_size, block_size, hidden_size = hidden_states.shape + normed = self._att_norm( + hidden_states.reshape(-1, hidden_size), + infer_state=infer_state, + layer_weight=layer_weight, + ).view(batch_size, block_size, hidden_size) + + q = layer_weight.q_proj.mm(normed.reshape(-1, hidden_size), use_custom_tensor_mananger=False) + kv = layer_weight.kv_proj.mm(normed.reshape(-1, hidden_size), use_custom_tensor_mananger=False) + + q = q.view(batch_size, block_size, self.tp_q_head_num_, self.head_dim_) + kv = kv.view(batch_size, block_size, self.tp_k_head_num_ + self.tp_v_head_num_, self.head_dim_) + k = kv[:, :, : self.tp_k_head_num_, :] + v = kv[:, :, self.tp_k_head_num_ :, :] + + q = layer_weight.q_norm_weight_( + input=q.reshape(-1, self.head_dim_), + eps=self.eps_, + alloc_func=torch.empty, + ).view(batch_size, block_size, self.tp_q_head_num_, self.head_dim_) + k = layer_weight.k_norm_weight_( + input=k.reshape(-1, self.head_dim_), + eps=self.eps_, + alloc_func=torch.empty, + ).view(batch_size, block_size, self.tp_k_head_num_, self.head_dim_) + + rotary_emb_fwd( + q.reshape(-1, self.tp_q_head_num_, self.head_dim_), + k.reshape(-1, self.tp_k_head_num_, self.head_dim_), + infer_state.position_cos, + infer_state.position_sin, + ) + cache_kv = torch.cat([k, v], dim=2).reshape( + batch_size * block_size, + self.tp_k_head_num_ + self.tp_v_head_num_, + self.head_dim_, + ) + return q, cache_kv diff --git a/lightllm/models/qwen3_dflash/layer_weights/__init__.py b/lightllm/models/qwen3_dflash/layer_weights/__init__.py new file mode 100644 index 0000000000..ffc51a2f3a --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_weights/__init__.py @@ -0,0 +1,11 @@ +from lightllm.models.qwen3_dflash.layer_weights.pre_and_post_layer_weight import ( + Qwen3DFlashPreAndPostLayerWeight, +) +from lightllm.models.qwen3_dflash.layer_weights.transformer_layer_weight import ( + Qwen3DFlashTransformerLayerWeight, +) + +__all__ = [ + "Qwen3DFlashPreAndPostLayerWeight", + "Qwen3DFlashTransformerLayerWeight", +] diff --git a/lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py new file mode 100644 index 0000000000..69807ceba4 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,63 @@ +from lightllm.common.basemodel import PreAndPostLayerWeight +from lightllm.common.basemodel.layer_weights.meta_weights import ( + EmbeddingWeight, + LMHeadWeight, + RMSNormWeight, + ROWMMWeight, +) +from lightllm.common.quantization import Quantcfg + + +class Qwen3DFlashPreAndPostLayerWeight(PreAndPostLayerWeight): + """Weights outside the DFlash decoder stack. + + DFlash checkpoints are stored as DSpark-family models, not as Qwen3 causal + LM checkpoints. Their top-level names are: + + - embed_tokens.weight: [vocab_size, hidden_size] + - fc.weight: [hidden_size, hidden_size * len(target_layer_ids)] + - hidden_norm.weight: [hidden_size] + - norm.weight: [hidden_size] + - lm_head.weight: [vocab_size, hidden_size] + """ + + def __init__(self, data_type, network_config, quant_cfg: Quantcfg): + super().__init__(data_type, network_config) + self.quant_cfg = quant_cfg + + hidden_size = network_config["hidden_size"] + vocab_size = network_config["vocab_size"] + target_layer_num = len(network_config["target_layer_ids"]) + + self.wte_weight_ = EmbeddingWeight( + dim=hidden_size, + vocab_size=vocab_size, + weight_name="embed_tokens.weight", + data_type=self.data_type_, + ) + self.fc_weight_ = ROWMMWeight( + in_dim=hidden_size * target_layer_num, + out_dims=[hidden_size], + weight_names="fc.weight", + data_type=self.data_type_, + quant_method=self.quant_cfg.get_quant_method(0, "fc"), + tp_rank=0, + tp_world_size=1, + ) + self.hidden_norm_weight_ = RMSNormWeight( + dim=hidden_size, + weight_name="hidden_norm.weight", + data_type=self.data_type_, + ) + self.final_norm_weight_ = RMSNormWeight( + dim=hidden_size, + weight_name="norm.weight", + data_type=self.data_type_, + ) + self.lm_head_weight_ = LMHeadWeight( + dim=hidden_size, + vocab_size=vocab_size, + weight_name="lm_head.weight", + data_type=self.data_type_, + ) + return diff --git a/lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py b/lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py new file mode 100644 index 0000000000..d7ae61044f --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py @@ -0,0 +1,102 @@ +from lightllm.common.basemodel.layer_weights.meta_weights import COLMMWeight, KVROWNMMWeight, RMSNormWeight, ROWMMWeight +from lightllm.models.llama.layer_weights.transformer_layer_weight import LlamaTransformerLayerWeight + + +class Qwen3DFlashTransformerLayerWeight(LlamaTransformerLayerWeight): + """DFlash decoder layer weights. + + The DSpark/DFlash Qwen3 layer uses Qwen3-style q/k RMSNorm, but its weight + prefix is `layers.{i}` rather than `model.layers.{i}`. + """ + + def _init_weight_names(self): + weight_prefix = f"layers.{self.layer_num_}" + self._q_weight_name = f"{weight_prefix}.self_attn.q_proj.weight" + self._q_bias_name = None + self._k_weight_name = f"{weight_prefix}.self_attn.k_proj.weight" + self._k_bias_name = None + self._v_weight_name = f"{weight_prefix}.self_attn.v_proj.weight" + self._v_bias_name = None + self._o_weight_name = f"{weight_prefix}.self_attn.o_proj.weight" + self._o_bias_name = None + + self._gate_weight_name = f"{weight_prefix}.mlp.gate_proj.weight" + self._gate_bias_name = None + self._up_weight_name = f"{weight_prefix}.mlp.up_proj.weight" + self._up_bias_name = None + self._down_weight_name = f"{weight_prefix}.mlp.down_proj.weight" + self._down_bias_name = None + + self._att_norm_weight_name = f"{weight_prefix}.input_layernorm.weight" + self._ffn_norm_weight_name = f"{weight_prefix}.post_attention_layernorm.weight" + self._q_norm_name = f"{weight_prefix}.self_attn.q_norm.weight" + self._k_norm_name = f"{weight_prefix}.self_attn.k_norm.weight" + return + + def _init_qkv(self): + in_dim = self.n_embed + q_out_dim = self.q_head_num_ * self.head_dim + self.q_proj = ROWMMWeight( + in_dim=in_dim, + out_dims=[q_out_dim], + weight_names=self._q_weight_name, + data_type=self.data_type_, + bias_names=self._q_bias_name, + quant_method=self.get_quant_method("q_proj"), + ) + self.kv_proj = KVROWNMMWeight( + in_dim=in_dim, + kv_head_num=self.k_head_num_, + head_dim=self.head_dim, + weight_names=[self._k_weight_name, self._v_weight_name], + data_type=self.data_type_, + bias_names=[self._k_bias_name, self._v_bias_name], + quant_method=self.get_quant_method("kv_proj"), + ) + return + + def _init_o(self): + in_dim = self.o_head_num_ * self.head_dim + out_dim = self.n_embed + self.o_proj = COLMMWeight( + in_dim=in_dim, + out_dims=[out_dim], + weight_names=self._o_weight_name, + data_type=self.data_type_, + bias_names=self._o_bias_name, + quant_method=self.get_quant_method("o_proj"), + ) + return + + def _init_ffn(self): + self.gate_up_proj = ROWMMWeight( + in_dim=self.n_embed, + out_dims=[self.n_inter, self.n_inter], + weight_names=[self._gate_weight_name, self._up_weight_name], + data_type=self.data_type_, + bias_names=[self._gate_bias_name, self._up_bias_name], + quant_method=self.get_quant_method("gate_up_proj"), + ) + self.down_proj = COLMMWeight( + in_dim=self.n_inter, + out_dims=[self.n_embed], + weight_names=self._down_weight_name, + data_type=self.data_type_, + bias_names=self._down_bias_name, + quant_method=self.get_quant_method("down_proj"), + ) + return + + def _init_norm(self): + super()._init_norm() + self.q_norm_weight_ = RMSNormWeight( + dim=self.head_dim, + weight_name=self._q_norm_name, + data_type=self.data_type_, + ) + self.k_norm_weight_ = RMSNormWeight( + dim=self.head_dim, + weight_name=self._k_norm_name, + data_type=self.data_type_, + ) + return diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py new file mode 100644 index 0000000000..cae7990bb4 --- /dev/null +++ b/lightllm/models/qwen3_dflash/model.py @@ -0,0 +1,159 @@ + +from lightllm.common.basemodel.attention import ( + BaseAttBackend, + Fa3AttBackend, + Fp8Fa3AttBackend, + get_decode_att_backend_class, + get_prefill_att_backend_class, +) +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.cuda_graph import CudaGraph +from lightllm.distributed.communication_op import dist_group_manager +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo +from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer +from lightllm.models.qwen3_dflash.layer_infer.pre_layer_infer import Qwen3DFlashPreLayerInfer +from lightllm.models.qwen3_dflash.layer_infer.transformer_layer_infer import Qwen3DFlashTransformerLayerInfer +from lightllm.models.qwen3_dflash.layer_weights.pre_and_post_layer_weight import Qwen3DFlashPreAndPostLayerWeight +from lightllm.models.qwen3_dflash.layer_weights.transformer_layer_weight import Qwen3DFlashTransformerLayerWeight + + +class Qwen3DFlashModel(LlamaTpPartModel): + """Qwen3 DFlash draft model. + + This is the LightLLM service port of the DeepSpec DFlash/DSpark Qwen3 model. + The service path enters through `forward(ModelInput)` using the same + primitive metadata as normal LightLLM prefill: `mem_indexes`, + `req_to_token_indexs`, sequence lengths, and prefill start locations. + + Target -> draft inputs: + - target hidden rows are committed into DFlash draft K/V. + - this model then materializes the next DFlash block from query/mask + embeddings and scratch KV slots. + + Draft output: + - logits are returned for the flattened [batch, draft_step] block rows. + The proposer maps them back to the standard + [verify_batch, draft_step + 1] speculative proposal shape. + + KV ownership: + - DFlash intentionally reuses `main_model.req_manager` and + `main_model.mem_manager` when the target/draft KV shape is compatible. + The target manager can be over-provisioned with draft layer slots and + extra token capacity so CPU cache, offload, and free logic remain unified. + - The DFlash-specific invariant is the layer/token-slot lifecycle: accepted + target hidden rows become committed draft K/V, while current block K/V is + scratch and must be released or overwritten after proposal/verification. + """ + + pre_and_post_weight_class = Qwen3DFlashPreAndPostLayerWeight + transformer_weight_class = Qwen3DFlashTransformerLayerWeight + pre_layer_infer_class = Qwen3DFlashPreLayerInfer + post_layer_infer_class = Qwen3DFlashPostLayerInfer + transformer_layer_infer_class = Qwen3DFlashTransformerLayerInfer + infer_state_class = Qwen3DFlashInferStateInfo + + def __init__(self, kvargs: dict): + self._pre_init(kvargs) + super().__init__(kvargs) + return + + def _pre_init(self, kvargs: dict): + self.main_model: TpPartBaseModel = kvargs.pop("main_model") + self.mtp_previous_draft_models = kvargs.pop("mtp_previous_draft_models") + kvargs["return_all_prompt_logics"] = True + return + + def _init_custom(self): + self._cos_cached = self.main_model._cos_cached + self._sin_cached = self.main_model._sin_cached + self.dist_group = dist_group_manager.get_default_group() + self.block_size = int(self.config["block_size"]) + self.mask_token_id = int(self.config["mask_token_id"]) + return + + def _init_req_manager(self): + self.req_manager = self.main_model.req_manager + return + + def _init_mem_manager(self): + # Intentionally shared with the target model. DFlash uses compatible + # KV shapes, so the main manager should be provisioned with the draft + # layer range and any extra temporary token capacity needed by block + # proposals. This keeps req/cpu-cache/offload/free paths unified. + self.mem_manager = self.main_model.mem_manager + return + + def _init_att_backend(self): + self.prefill_att_backend: BaseAttBackend = get_prefill_att_backend_class(index=0)(model=self) + try: + self.decode_att_backend: BaseAttBackend = get_decode_att_backend_class( + index=0, + priority_list=["fa3"], + )(model=self) + except KeyError as exc: + raise NotImplementedError( + "Qwen3DFlashModel requires FA3 decode attention: " + "block draft attention is non-causal and Triton/FlashInfer decode paths do not honor decode_causal." + ) from exc + if not isinstance(self.decode_att_backend, (Fa3AttBackend, Fp8Fa3AttBackend)): + raise NotImplementedError( + "Qwen3DFlashModel requires FA3 decode attention: " + "block draft attention is non-causal and Triton/FlashInfer decode paths do not honor decode_causal." + ) + return + + def _init_infer_layer(self, start_layer_index=None): + assert start_layer_index is None + self.draft_layer_start = len(self.main_model.layers_infer) + self.draft_layer_start += sum( + len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models + ) + super()._init_infer_layer(start_layer_index=self.draft_layer_start) + return + + def _init_weights(self, start_layer_index=None): + assert start_layer_index is None + self.pre_post_weight = self.pre_and_post_weight_class( + self.data_type, + network_config=self.config, + quant_cfg=self.quant_cfg, + ) + self.trans_layers_weight = [ + self.transformer_weight_class( + i, + self.data_type, + network_config=self.config, + quant_cfg=self.quant_cfg, + ) + for i in range(self.config["n_layer"]) + ] + return + + def _gen_special_model_input(self, token_num: int): + return {"mtp_draft_input_hiddens": None} + + def _autotune_warmup(self): + return + + def _init_padded_req(self): + return + + def _init_cudagraph(self): + if self.disable_cudagraph or self.args.enable_decode_microbatch_overlap: + self.graph = None + return + + self.graph = CudaGraph( + max_batch_size=self.graph_max_batch_size, + max_len_in_batch=self.graph_max_len_in_batch, + tp_world_size=self.tp_world_size_, + ) + return + + def _init_prefill_cuda_graph(self): + self.prefill_graph = None + return + + def _check_max_len_infer(self): + return diff --git a/lightllm/models/qwen3_dspark/__init__.py b/lightllm/models/qwen3_dspark/__init__.py new file mode 100644 index 0000000000..c59c5bd2de --- /dev/null +++ b/lightllm/models/qwen3_dspark/__init__.py @@ -0,0 +1,3 @@ +from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel + +__all__ = ["Qwen3DSparkModel"] diff --git a/lightllm/models/qwen3_dspark/layer_infer/__init__.py b/lightllm/models/qwen3_dspark/layer_infer/__init__.py new file mode 100644 index 0000000000..4e412110b5 --- /dev/null +++ b/lightllm/models/qwen3_dspark/layer_infer/__init__.py @@ -0,0 +1,4 @@ +from .post_layer_infer import Qwen3DSparkPostLayerInfer + + +__all__ = ["Qwen3DSparkPostLayerInfer"] diff --git a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py new file mode 100644 index 0000000000..e5efa6768f --- /dev/null +++ b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py @@ -0,0 +1,225 @@ +import numpy as np +import torch +import torch.nn.functional as F + +from lightllm.distributed.communication_op import all_gather +from lightllm.models.llama.infer_struct import LlamaInferStateInfo +from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer +from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import ( + Qwen3DSparkPreAndPostLayerWeight, +) + + +class Qwen3DSparkPostLayerInfer(Qwen3DFlashPostLayerInfer): + """DSpark post layer. + + The block backbone produces one flat logits row per block position. This + post layer applies DSpark's sequential Markov correction before returning + logits, and stores raw confidence logits for dynamic verify scheduling. + """ + + def __init__(self, network_config): + super().__init__(network_config) + self.block_size_ = int(network_config["block_size"]) + self.markov_rank_ = int(network_config.get("markov_rank", 0)) + self.markov_head_type_ = str(network_config.get("markov_head_type", "")).lower() + self.enable_confidence_head_ = bool(network_config.get("enable_confidence_head", False)) + self.confidence_head_with_markov_ = bool(network_config.get("confidence_head_with_markov", False)) + self.mtp_draft_confidence_logits = None + + def pop_mtp_draft_confidence_logits(self): + logits = self.mtp_draft_confidence_logits + self.mtp_draft_confidence_logits = None + return logits + + def has_markov_head(self) -> bool: + return self.markov_rank_ > 0 + + def has_confidence_head(self) -> bool: + return self.enable_confidence_head_ + + def _linear_parameter(self, input_tensor: torch.Tensor, parameter_weight) -> torch.Tensor: + assert parameter_weight is not None + weight = parameter_weight.weight + bias = parameter_weight.bias + return F.linear(input_tensor.to(dtype=weight.dtype), weight, bias) + + def _markov_prev_embeddings( + self, + token_ids: torch.Tensor, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ) -> torch.Tensor: + assert layer_weight.markov_w1_weight_ is not None + return F.embedding(token_ids.long(), layer_weight.markov_w1_weight_.weight) + + def _markov_project_bias( + self, + latent_states: torch.Tensor, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ) -> torch.Tensor: + assert layer_weight.markov_w2_weight_ is not None + weight = layer_weight.markov_w2_weight_.weight + return F.linear(latent_states.to(dtype=weight.dtype), weight) + + def _markov_step_bias( + self, + *, + prev_token_ids: torch.Tensor, + hidden_states: torch.Tensor, + state: torch.Tensor, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ): + prev_embeddings = self._markov_prev_embeddings(prev_token_ids, layer_weight) + if self.markov_head_type_ == "vanilla": + return state, self._markov_project_bias(prev_embeddings, layer_weight) + + hidden_states = hidden_states.to(dtype=prev_embeddings.dtype) + if self.markov_head_type_ == "gated": + gate_input = torch.cat([hidden_states, prev_embeddings], dim=-1) + gate = torch.sigmoid(self._linear_parameter(gate_input, layer_weight.markov_gate_proj_weight_)) + return state, self._markov_project_bias(gate * prev_embeddings, layer_weight) + + assert self.markov_head_type_ == "rnn" + if state is None: + state = torch.zeros_like(prev_embeddings) + joint_input = torch.cat([state, prev_embeddings, hidden_states], dim=-1) + joint = self._linear_parameter(joint_input, layer_weight.markov_joint_proj_weight_) + gate_raw, candidate_raw, output_raw = joint.chunk(3, dim=-1) + gate = torch.sigmoid(gate_raw) + candidate = torch.tanh(candidate_raw) + state = gate * state + (1.0 - gate) * candidate + return state, self._markov_project_bias(torch.tanh(output_raw), layer_weight) + + @torch.no_grad() + def apply_markov_logits( + self, + base_logits: torch.Tensor, + *, + block_hidden: torch.Tensor, + anchor_token_ids: torch.Tensor, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ): + if not self.has_markov_head(): + return base_logits, torch.argmax(base_logits, dim=-1) + + sampled_tokens = [] + corrected_logits = [] + prev_token_ids = anchor_token_ids.long() + state = None + for step_idx in range(base_logits.shape[1]): + state, markov_bias = self._markov_step_bias( + prev_token_ids=prev_token_ids, + hidden_states=block_hidden[:, step_idx, :], + state=state, + layer_weight=layer_weight, + ) + step_logits = base_logits[:, step_idx, :] + markov_bias + next_token_ids = torch.argmax(step_logits, dim=-1) + sampled_tokens.append(next_token_ids) + corrected_logits.append(step_logits.unsqueeze(1)) + prev_token_ids = next_token_ids + + return torch.cat(corrected_logits, dim=1), torch.stack(sampled_tokens, dim=1) + + @torch.no_grad() + def predict_confidence_logits( + self, + block_hidden: torch.Tensor, + *, + anchor_token_ids: torch.Tensor, + sampled_tokens: torch.Tensor, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ): + if not self.has_confidence_head(): + return None + + features = block_hidden + if self.confidence_head_with_markov_: + prev_token_ids = torch.cat( + [anchor_token_ids.view(-1, 1), sampled_tokens[:, :-1]], + dim=1, + ) + prev_embeddings = self._markov_prev_embeddings(prev_token_ids, layer_weight).to(dtype=block_hidden.dtype) + features = torch.cat([block_hidden, prev_embeddings], dim=-1) + + logits = self._linear_parameter(features, layer_weight.confidence_head_weight_) + return logits.float().squeeze(-1) + + def _token_forward_with_hidden( + self, + input_embdings: torch.Tensor, + infer_state: LlamaInferStateInfo, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ): + last_input, token_num = self._slice_get_last_input(input_embdings, infer_state) + input_embdings_dtype = input_embdings.dtype + head_hidden = last_input + normed_input = self._norm(last_input, infer_state, layer_weight) + lm_head_input = normed_input.permute(1, 0).reshape(-1, token_num) + logic_batch = layer_weight.lm_head_weight_(input=lm_head_input, alloc_func=self.alloc_tensor) + normed_input = None + lm_head_input = None + vocab_size = layer_weight.lm_head_weight_.vocab_size + if self.tp_world_size_ == 1: + gather_data = logic_batch + else: + gather_data = self.alloc_tensor((vocab_size, token_num), dtype=input_embdings_dtype) + split_indexes = np.linspace(0, vocab_size, self.tp_world_size_ + 1, dtype=np.int64) + all_gather( + [gather_data[split_indexes[i] : split_indexes[i + 1], :] for i in range(self.tp_world_size_)], + logic_batch, + group=infer_state.dist_group, + async_op=False, + ) + logic_batch = None + logits = self.alloc_tensor( + (token_num, vocab_size), + dtype=torch.float32, + ) + logits[:, :] = gather_data.permute(1, 0) + gather_data = None + return logits, head_hidden + + def token_forward( + self, + input_embdings: torch.Tensor, + infer_state: LlamaInferStateInfo, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ): + self.mtp_draft_confidence_logits = None + if self._is_commit_prefill(infer_state): + return super().token_forward( + input_embdings=input_embdings, + infer_state=infer_state, + layer_weight=layer_weight, + ) + + logits, head_hidden = self._token_forward_with_hidden( + input_embdings=input_embdings, + infer_state=infer_state, + layer_weight=layer_weight, + ) + if infer_state.is_prefill: + return logits + + assert ( + logits.shape[0] % self.block_size_ == 0 + ), f"DSpark draft logits rows must be a multiple of block_size={self.block_size_}, got {logits.shape[0]}" + num_reqs = logits.shape[0] // self.block_size_ + block_logits = logits.reshape(num_reqs, self.block_size_, -1) + block_hidden = head_hidden.reshape(num_reqs, self.block_size_, -1) + anchor_token_ids = infer_state.input_ids.reshape(num_reqs, self.block_size_)[:, 0] + + corrected_logits, sampled_tokens = self.apply_markov_logits( + block_logits, + block_hidden=block_hidden, + anchor_token_ids=anchor_token_ids, + layer_weight=layer_weight, + ) + self.mtp_draft_confidence_logits = self.predict_confidence_logits( + block_hidden, + anchor_token_ids=anchor_token_ids, + sampled_tokens=sampled_tokens, + layer_weight=layer_weight, + ) + return corrected_logits.reshape(logits.shape[0], -1) diff --git a/lightllm/models/qwen3_dspark/layer_weights/__init__.py b/lightllm/models/qwen3_dspark/layer_weights/__init__.py new file mode 100644 index 0000000000..7c9e42085a --- /dev/null +++ b/lightllm/models/qwen3_dspark/layer_weights/__init__.py @@ -0,0 +1,5 @@ +from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import ( + Qwen3DSparkPreAndPostLayerWeight, +) + +__all__ = ["Qwen3DSparkPreAndPostLayerWeight"] diff --git a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py new file mode 100644 index 0000000000..3555ad1405 --- /dev/null +++ b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,67 @@ +from lightllm.common.basemodel.layer_weights.meta_weights import ParameterWeight +from lightllm.common.quantization import Quantcfg +from lightllm.models.qwen3_dflash.layer_weights.pre_and_post_layer_weight import Qwen3DFlashPreAndPostLayerWeight + + +class Qwen3DSparkPreAndPostLayerWeight(Qwen3DFlashPreAndPostLayerWeight): + """DSpark heads on top of the shared DFlash block backbone weights.""" + + def __init__(self, data_type, network_config, quant_cfg: Quantcfg): + super().__init__(data_type, network_config, quant_cfg) + + hidden_size = network_config["hidden_size"] + vocab_size = network_config["vocab_size"] + markov_rank = int(network_config.get("markov_rank", 0)) + enable_confidence_head = bool(network_config.get("enable_confidence_head", False)) + confidence_head_with_markov = bool(network_config.get("confidence_head_with_markov", False)) + assert ( + not confidence_head_with_markov or markov_rank > 0 + ), "confidence_head_with_markov requires markov_rank > 0" + + self.markov_w1_weight_ = None + self.markov_w2_weight_ = None + self.markov_gate_proj_weight_ = None + self.markov_joint_proj_weight_ = None + self.markov_rank = markov_rank + self.markov_head_type = str(network_config.get("markov_head_type", "")).lower() + if markov_rank > 0: + self.markov_w1_weight_ = ParameterWeight( + weight_name="markov_head.markov_w1.weight", + data_type=self.data_type_, + weight_shape=(vocab_size, markov_rank), + ) + self.markov_w2_weight_ = ParameterWeight( + weight_name="markov_head.markov_w2.weight", + data_type=self.data_type_, + weight_shape=(vocab_size, markov_rank), + ) + if self.markov_head_type == "gated": + self.markov_gate_proj_weight_ = ParameterWeight( + weight_name="markov_head.gate_proj.weight", + bias_name="markov_head.gate_proj.bias", + data_type=self.data_type_, + weight_shape=(markov_rank, hidden_size + markov_rank), + bias_shape=(markov_rank,), + ) + elif self.markov_head_type == "rnn": + self.markov_joint_proj_weight_ = ParameterWeight( + weight_name="markov_head.joint_proj.weight", + bias_name="markov_head.joint_proj.bias", + data_type=self.data_type_, + weight_shape=(3 * markov_rank, hidden_size + 2 * markov_rank), + bias_shape=(3 * markov_rank,), + ) + else: + assert self.markov_head_type == "vanilla", f"unsupported DSpark markov head {self.markov_head_type}" + + self.confidence_head_weight_ = None + if enable_confidence_head: + confidence_input_dim = hidden_size + (markov_rank if confidence_head_with_markov else 0) + self.confidence_head_weight_ = ParameterWeight( + weight_name="confidence_head.proj.weight", + bias_name="confidence_head.proj.bias", + data_type=self.data_type_, + weight_shape=(1, confidence_input_dim), + bias_shape=(1,), + ) + return diff --git a/lightllm/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py new file mode 100644 index 0000000000..99029c5a1a --- /dev/null +++ b/lightllm/models/qwen3_dspark/model.py @@ -0,0 +1,15 @@ +from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel +from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer +from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import Qwen3DSparkPreAndPostLayerWeight + + +class Qwen3DSparkModel(Qwen3DFlashModel): + """Qwen3 DSpark draft model. + + DSpark keeps the DFlash block backbone. Its extra Markov/confidence heads + run in the post layer so the proposer can consume corrected logits through + the same token-generation path as DFlash. + """ + + pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight + post_layer_infer_class = Qwen3DSparkPostLayerInfer diff --git a/lightllm/models/qwen3_eagle/__init__.py b/lightllm/models/qwen3_eagle/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/qwen3_eagle/layer_infer/__init__.py b/lightllm/models/qwen3_eagle/layer_infer/__init__.py new file mode 100644 index 0000000000..9efa61f978 --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_infer/__init__.py @@ -0,0 +1,7 @@ +from lightllm.models.qwen3_eagle.layer_infer.pre_layer_infer import Qwen3EaglePreLayerInfer +from lightllm.models.qwen3_eagle.layer_infer.transformer_layer_infer import Qwen3EagleTransformerLayerInfer + +__all__ = [ + "Qwen3EaglePreLayerInfer", + "Qwen3EagleTransformerLayerInfer", +] diff --git a/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py new file mode 100644 index 0000000000..9fed1ba284 --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py @@ -0,0 +1,55 @@ +from lightllm.common.basemodel.infer_struct import InferStateInfo +from lightllm.models.llama.layer_infer.pre_layer_infer import LlamaPreLayerInfer +from lightllm.models.qwen3_eagle.layer_weights.pre_and_post_layer_weight import Qwen3EaglePreAndPostLayerWeight + + +class Qwen3EaglePreLayerInfer(LlamaPreLayerInfer): + """Eagle3 draft-token embedding plus target-hidden projection.""" + + def __init__(self, network_config): + super().__init__(network_config) + self.hidden_size_ = network_config["hidden_size"] + return + + def prepare_mtp_draft_hiddens( + self, + infer_state: InferStateInfo, + layer_weight: Qwen3EaglePreAndPostLayerWeight, + ) -> None: + # Keep the ModelInput hidden raw for CUDA graph replay; Eagle layers consume this working buffer. + infer_state.eagle_draft_hidden_states = self.project_mtp_draft_hiddens( + infer_state.mtp_draft_input_hiddens, + layer_weight, + ) + return + + def project_mtp_draft_hiddens( + self, + target_hiddens, + layer_weight: Qwen3EaglePreAndPostLayerWeight, + use_custom_tensor_mananger: bool = True, + ): + if target_hiddens is None or target_hiddens.shape[-1] == self.hidden_size_: + return target_hiddens + return layer_weight.fc_weight_.mm( + target_hiddens, + use_custom_tensor_mananger=use_custom_tensor_mananger, + ) + + def context_forward( + self, + input_ids, + infer_state: InferStateInfo, + layer_weight: Qwen3EaglePreAndPostLayerWeight, + ): + self.prepare_mtp_draft_hiddens(infer_state, layer_weight) + return super().context_forward(input_ids, infer_state, layer_weight) + + def token_forward( + self, + input_ids, + infer_state: InferStateInfo, + layer_weight: Qwen3EaglePreAndPostLayerWeight, + ): + self.prepare_mtp_draft_hiddens(infer_state, layer_weight) + return super().token_forward(input_ids, infer_state, layer_weight) diff --git a/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py new file mode 100644 index 0000000000..946a6f8b50 --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py @@ -0,0 +1,67 @@ +import torch + +from lightllm.common.basemodel.infer_struct import InferStateInfo +from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer +from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd +from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight + + +class Qwen3EagleTransformerLayerInfer(LlamaTransformerLayerInfer): + def __init__(self, layer_num, network_config): + super().__init__(layer_num, network_config) + self.head_dim_ = network_config["head_dim"] + return + + def _get_qkv(self, input, infer_state: InferStateInfo, layer_weight: Qwen3EagleTransformerLayerWeight): + input_part = self._att_norm(input, infer_state, layer_weight) + target_part = layer_weight.hidden_norm_weight_._native_forward( + input=infer_state.eagle_draft_hidden_states, + eps=self.eps_, + alloc_func=self.alloc_tensor, + ) + input = torch.cat([input_part, target_part], dim=-1) + input = input.view(-1, self.embed_dim_ * 2) + input = self._tpsp_allgather(input, infer_state) + q = layer_weight.q_proj.mm(input) + cache_kv = layer_weight.kv_proj.mm(input) + if layer_weight.qk_norm_weight_ is not None: + layer_weight.qk_norm_weight_( + q, + cache_kv[:, : self.tp_k_head_num_ * self.head_dim_], + eps=self.eps_, + ) + cache_kv = cache_kv.view(-1, (self.tp_k_head_num_ + self.tp_v_head_num_), self.head_dim_) + rotary_emb_fwd( + q.view(-1, self.tp_q_head_num_, self.head_dim_), + cache_kv[:, : self.tp_k_head_num_, :], + infer_state.position_cos, + infer_state.position_sin, + ) + + if infer_state.need_dp_prefill_balance: + q = infer_state._all_to_all_unbalance_get(data=q) + cache_kv = infer_state._all_to_all_unbalance_get(data=cache_kv) + + return q, cache_kv + + def context_forward(self, input_embdings, infer_state: InferStateInfo, layer_weight): + o = self.context_attention_forward(input_embdings, infer_state, layer_weight) + + hidden_states = infer_state.eagle_draft_hidden_states + o.view(-1, self.embed_dim_) + ffn_input = self._ffn_norm(hidden_states, infer_state, layer_weight) + ffn_out = self._ffn(ffn_input, infer_state, layer_weight) + + hidden_states = hidden_states + ffn_out.view(-1, self.embed_dim_) + infer_state.eagle_draft_hidden_states = hidden_states + return hidden_states + + def token_forward(self, input_embdings, infer_state: InferStateInfo, layer_weight): + o = self.token_attention_forward(input_embdings, infer_state, layer_weight) + + hidden_states = infer_state.eagle_draft_hidden_states + o.view(-1, self.embed_dim_) + ffn_input = self._ffn_norm(hidden_states, infer_state, layer_weight) + ffn_out = self._ffn(ffn_input, infer_state, layer_weight) + + hidden_states = hidden_states + ffn_out.view(-1, self.embed_dim_) + infer_state.eagle_draft_hidden_states = hidden_states + return hidden_states diff --git a/lightllm/models/qwen3_eagle/layer_weights/__init__.py b/lightllm/models/qwen3_eagle/layer_weights/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py new file mode 100644 index 0000000000..979d6c04d2 --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,64 @@ +import torch +from lightllm.common.basemodel import PreAndPostLayerWeight +from lightllm.common.basemodel.layer_weights.meta_weights import ( + EmbeddingWeight, + LMHeadWeight, + ParameterWeight, + RMSNormWeight, + ROWMMWeight, +) +from lightllm.common.quantization import Quantcfg + + +class Qwen3EaglePreAndPostLayerWeight(PreAndPostLayerWeight): + def __init__(self, data_type, network_config, quant_cfg: Quantcfg): + super().__init__(data_type, network_config) + self.quant_cfg: Quantcfg = quant_cfg + hidden_size = network_config["hidden_size"] + target_layer_num = len(network_config.get("target_layer_ids", [0, 1, 2])) + draft_vocab_size = network_config.get("draft_vocab_size") + target_vocab_size = network_config.get("target_vocab_size") + vocab_size = draft_vocab_size if draft_vocab_size is not None else network_config["vocab_size"] + self.fc_weight_ = ROWMMWeight( + in_dim=hidden_size * target_layer_num, + out_dims=[hidden_size], + weight_names="fc.weight", + quant_method=self.quant_cfg.get_quant_method(0, "fc"), + data_type=self.data_type_, + tp_rank=0, + tp_world_size=1, + ) + self.final_norm_weight_: RMSNormWeight = RMSNormWeight( + dim=hidden_size, + weight_name="norm.weight", + data_type=self.data_type_, + ) + + self.d2t_weight_ = None + self.t2d_weight_ = None + if draft_vocab_size is not None and target_vocab_size is not None: + self.d2t_weight_ = ParameterWeight( + weight_name="d2t", + data_type=torch.int64, + weight_shape=(draft_vocab_size,), + ) + self.t2d_weight_ = ParameterWeight( + weight_name="t2d", + data_type=torch.bool, + weight_shape=(target_vocab_size,), + ) + + self.lm_head_weight_: LMHeadWeight = LMHeadWeight( + dim=hidden_size, + vocab_size=vocab_size, + weight_name="lm_head.weight", + data_type=self.data_type_, + ) + self.wte_weight_ = EmbeddingWeight( + dim=hidden_size, + vocab_size=vocab_size, + weight_name="embed_tokens.weight", + data_type=self.data_type_, + ) + + return diff --git a/lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py b/lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py new file mode 100644 index 0000000000..26d2ea3345 --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py @@ -0,0 +1,111 @@ +from lightllm.common.basemodel.layer_weights.meta_weights.mm_weight.colmm_weight import COLMMWeight +from lightllm.common.basemodel.layer_weights.meta_weights.mm_weight.rowmm_weight import KVROWNMMWeight, ROWMMWeight +from lightllm.common.basemodel.layer_weights.meta_weights.norm_weight import QKRMSNORMWeight, RMSNormWeight +from lightllm.models.llama.layer_weights.transformer_layer_weight import LlamaTransformerLayerWeight + +""" +midlayer.hidden_norm.weight [2048] +midlayer.input_layernorm.weight [2048] +midlayer.mlp.down_proj.weight [2048, 6144] +midlayer.mlp.gate_proj.weight [6144, 2048] +midlayer.mlp.up_proj.weight [6144, 2048] +midlayer.post_attention_layernorm.weight [2048] +midlayer.self_attn.k_proj.weight [512, 4096] +midlayer.self_attn.o_proj.weight [2048, 4096] +midlayer.self_attn.q_proj.weight [4096, 4096] +midlayer.self_attn.v_proj.weight [512, 4096] +""" + + +class Qwen3EagleTransformerLayerWeight(LlamaTransformerLayerWeight): + def _init_weight_names(self): + super()._init_weight_names() + if self.network_config_["architectures"][0] in ["Eagle3Speculator", "Qwen3Eagle3Model"]: + weight_prefix = f"layers.{self.layer_num_}" + else: + weight_prefix = "midlayer" + self._q_weight_name = f"{weight_prefix}.self_attn.q_proj.weight" + self._k_weight_name = f"{weight_prefix}.self_attn.k_proj.weight" + self._v_weight_name = f"{weight_prefix}.self_attn.v_proj.weight" + self._kv_weight_name = f"{weight_prefix}.self_attn.kv_proj.weight" + self._o_weight_name = f"{weight_prefix}.self_attn.o_proj.weight" + + self._gate_weight_name = f"{weight_prefix}.mlp.gate_proj.weight" + self._up_weight_name = f"{weight_prefix}.mlp.up_proj.weight" + self._down_weight_name = f"{weight_prefix}.mlp.down_proj.weight" + self._gate_up_bias_name = None + + self._att_norm_weight_name = f"{weight_prefix}.input_layernorm.weight" + self._ffn_norm_weight_name = f"{weight_prefix}.post_attention_layernorm.weight" + self._hidden_norm_weight_name = f"{weight_prefix}.hidden_norm.weight" + self._q_norm_name = f"{weight_prefix}.self_attn.q_norm.weight" + self._k_norm_name = f"{weight_prefix}.self_attn.k_norm.weight" + + def _init_qkv(self): + in_dim = self.n_embed * 2 + q_out_dim = self.q_head_num_ * self.head_dim + self.q_proj = ROWMMWeight( + in_dim=in_dim, + out_dims=[q_out_dim], + weight_names=self._q_weight_name, + data_type=self.data_type_, + bias_names=self._q_bias_name, + quant_method=self.get_quant_method("q_proj"), + ) + self.kv_proj = KVROWNMMWeight( + in_dim=in_dim, + kv_head_num=self.k_head_num_, + head_dim=self.head_dim, + weight_names=[self._k_weight_name, self._v_weight_name], + data_type=self.data_type_, + bias_names=[self._k_bias_name, self._v_bias_name], + quant_method=self.get_quant_method("kv_proj"), + ) + + def _init_o(self): + in_dim = self.o_head_num_ * self.head_dim + out_dim = self.n_embed + self.o_proj = COLMMWeight( + in_dim=in_dim, + out_dims=[out_dim], + weight_names=self._o_weight_name, + data_type=self.data_type_, + bias_names=self._o_bias_name, + quant_method=self.get_quant_method("o_proj"), + ) + + def _init_ffn(self): + self.gate_up_proj = ROWMMWeight( + in_dim=self.n_embed, + out_dims=[self.n_inter, self.n_inter], + weight_names=[self._gate_weight_name, self._up_weight_name], + data_type=self.data_type_, + bias_names=[self._gate_bias_name, self._up_bias_name], + quant_method=self.get_quant_method("gate_up_proj"), + ) + self.down_proj = COLMMWeight( + in_dim=self.n_inter, + out_dims=[self.n_embed], + weight_names=self._down_weight_name, + data_type=self.data_type_, + bias_names=self._down_bias_name, + quant_method=self.get_quant_method("down_proj"), + ) + + def _init_norm(self): + super()._init_norm() + hidden_size = self.network_config_["hidden_size"] + self.hidden_norm_weight_ = RMSNormWeight( + dim=hidden_size, + weight_name=self._hidden_norm_weight_name, + data_type=self.data_type_, + ) + self.qk_norm_weight_ = None + architecture = (self.network_config_.get("architectures") or [""])[0] + if architecture in {"Eagle3Speculator", "Qwen3Eagle3Model"}: + self.qk_norm_weight_ = QKRMSNORMWeight( + dim=self.head_dim, + q_weight_name=self._q_norm_name, + k_weight_name=self._k_norm_name, + data_type=self.data_type_, + ) diff --git a/lightllm/models/qwen3_eagle/model.py b/lightllm/models/qwen3_eagle/model.py new file mode 100644 index 0000000000..e6d012da37 --- /dev/null +++ b/lightllm/models/qwen3_eagle/model.py @@ -0,0 +1,85 @@ +import torch +from typing import List + +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.qwen3_eagle.layer_infer.pre_layer_infer import Qwen3EaglePreLayerInfer +from lightllm.models.qwen3_eagle.layer_infer.transformer_layer_infer import Qwen3EagleTransformerLayerInfer +from lightllm.models.qwen3_eagle.layer_weights.pre_and_post_layer_weight import Qwen3EaglePreAndPostLayerWeight +from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight + + +class Qwen3EagleModel(LlamaTpPartModel): + pre_and_post_weight_class = Qwen3EaglePreAndPostLayerWeight + pre_layer_infer_class = Qwen3EaglePreLayerInfer + + transformer_weight_class = Qwen3EagleTransformerLayerWeight + transformer_layer_infer_class = Qwen3EagleTransformerLayerInfer + + def __init__(self, kvargs: dict): + self._pre_init(kvargs) + super().__init__(kvargs) + return + + def _pre_init(self, kvargs: dict): + self.main_model: TpPartBaseModel = kvargs.pop("main_model") + self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") + return + + def _gen_special_model_input(self, token_num: int): + return self._gen_mtp_draft_special_model_input(token_num) + + def _init_custom(self): + self._cos_cached = self.main_model._cos_cached + self._sin_cached = self.main_model._sin_cached + return + + def _init_req_manager(self): + self.req_manager = self.main_model.req_manager + return + + def _init_mem_manager(self): + self.mem_manager = self.main_model.mem_manager + return + + def _init_weights(self, start_layer_index=None): + assert start_layer_index is None + self.pre_post_weight = self.pre_and_post_weight_class( + self.data_type, network_config=self.config, quant_cfg=self.quant_cfg + ) + # Compressed-vocabulary EAGLE3 checkpoints store a draft lm_head and + # vocabulary maps, but their input ids still belong to the target + # vocabulary. Reuse the target embedding instead of requiring a draft + # embed_tokens.weight that these checkpoints do not contain. + if self.config.get("draft_vocab_size") is not None: + target_embedding = getattr(self.main_model.pre_post_weight, "wte_weight_", None) + assert target_embedding is not None, "compressed-vocab EAGLE3 requires target token embeddings" + self.pre_post_weight.wte_weight_ = target_embedding + self.trans_layers_weight = [ + self.transformer_weight_class( + i, + self.data_type, + network_config=self.config, + quant_cfg=self.quant_cfg, + ) + for i in range(self.config["n_layer"]) + ] + return + + def _init_infer_layer(self, start_layer_index=None): + assert start_layer_index is None + total_pre_layers_num = len(self.main_model.layers_infer) + total_pre_layers_num += sum( + len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models + ) + super()._init_infer_layer(start_layer_index=total_pre_layers_num) + return + + # d2t stores per-token offsets: target_id = draft_id + d2t[draft_id]. + @torch.no_grad() + def map_draft_vocab_to_main_vocab(self, draft_token_ids: torch.Tensor) -> torch.Tensor: + if self.pre_post_weight.d2t_weight_ is not None: + draft_token_ids = draft_token_ids + self.pre_post_weight.d2t_weight_.weight[draft_token_ids].to( + dtype=draft_token_ids.dtype + ) + return draft_token_ids diff --git a/lightllm/models/qwen3_moe_mtp/model.py b/lightllm/models/qwen3_moe_mtp/model.py index d9854250e2..f91d9a47e7 100644 --- a/lightllm/models/qwen3_moe_mtp/model.py +++ b/lightllm/models/qwen3_moe_mtp/model.py @@ -8,10 +8,6 @@ class Qwen3MOEMTPModel(Qwen3MOEModel): - - # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). - is_mtp_draft_model = True - pre_and_post_weight_class = Qwen3MOEMTPPreAndPostLayerWeight pre_layer_infer_class = Deepseek3MTPPreLayerInfer @@ -28,6 +24,9 @@ def _pre_init(self, kvargs: dict): self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") return + def _gen_special_model_input(self, token_num: int): + return self._gen_mtp_draft_special_model_input(token_num) + def _init_custom(self): self._cos_cached = self.main_model._cos_cached self._sin_cached = self.main_model._sin_cached diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 5838b0d0da..3526980efc 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -613,6 +613,34 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: a new CUDA graph will be generated for every increment of graph_grow_step_size. """, ) + parser.add_argument( + "--mtp_draft_graph_max_batch_size", + type=int, + default=None, + help=""" + Optional logical CUDA graph batch-size limit for MTP draft models. + Defaults to graph_max_batch_size. The value is expanded by the draft + model's decode graph group size in the same way as the main setting. + """, + ) + parser.add_argument( + "--mtp_draft_graph_split_batch_size", + type=int, + default=None, + help=""" + Optional dense-prefix CUDA graph limit for MTP draft models. + Defaults to graph_split_batch_size. + """, + ) + parser.add_argument( + "--mtp_draft_graph_grow_step_size", + type=int, + default=None, + help=""" + Optional CUDA graph batch-size growth step for MTP draft models. + Defaults to graph_grow_step_size. + """, + ) parser.add_argument( "--graph_max_len_in_batch", type=int, @@ -733,13 +761,17 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "eagle_with_att", "vanilla_no_att", "eagle_no_att", + "eagle3", + "qwen3next_vanilla", + "qwen3next_eagle", + "dspark", + "dflash", None, ], default=None, - help="""Supported MTP modes. - None: Disables MTP. - *_with_att: Uses the MTP model with an attention mechanism to predict the next draft token. - *_no_att: Uses the MTP model without an attention module to predict the next draft token.""", + help="""Speculative decoding mode. + *_with_att and *_no_att select attention or non-attention MTP drafts; + eagle3 uses a recurrent EAGLE3 draft; dspark and dflash use block draft models.""", ) parser.add_argument( "--mtp_draft_model_dir", @@ -753,11 +785,13 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--mtp_step", type=int, default=0, - help="""Specifies the number of additional tokens to predict using the draft model. - Currently, this feature supports only the DeepSeekV3 model. - Increasing this value allows for more predictions, - but ensure that the model is compatible with the specified step count. - currently, deepseekv3 model only support 1 step""", + help="""Number of additional draft tokens per request. + For DSpark and DFlash this value is derived from the draft checkpoint block_size.""", + ) + parser.add_argument( + "--mtp_dynamic_verify", + action="store_true", + help="""Whether to enable dynamic verification for MTP multi-prediction results.""", ) parser.add_argument( "--kv_quant_calibration_config_path", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 67ae286a7c..5c4879060f 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -3,6 +3,8 @@ import uuid import subprocess import math +from dataclasses import replace +from transformers.configuration_utils import PretrainedConfig from lightllm.utils.start_utils import process_manager from .metrics.manager import start_metric_manager from .embed_cache.manager import start_cache_manager @@ -26,10 +28,37 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args +from lightllm.common.speculative import SpeculativeConfig, get_dspark_family_block_size logger = init_logger(__name__) +def normalize_block_mtp_step_from_first_draft_config( + args: StartArgs, spec_config: SpeculativeConfig +) -> SpeculativeConfig: + if not spec_config.uses_block_draft_model: + return spec_config + + assert args.mtp_draft_model_dir is not None and len(args.mtp_draft_model_dir) > 0 + mtp_model_cfg, _ = PretrainedConfig.get_config_dict(args.mtp_draft_model_dir[0]) + block_size = get_dspark_family_block_size( + mtp_model_cfg, + require_confidence_head=spec_config.is_dspark, + ) + configured_step = int(args.mtp_step) + if configured_step not in (0, block_size): + logger.warning( + "Overriding mtp_step=%s with block draft config block_size=%s for %s mode", + configured_step, + block_size, + spec_config.mode, + ) + args.mtp_step = block_size + spec_config = replace(spec_config, step=block_size) + spec_config.validate() + return spec_config + + def _set_envs_and_config(args: StartArgs): mp.set_start_method("spawn", force=True) @@ -97,9 +126,9 @@ def _launch_subprocesses(args: StartArgs): args.embed_cache_storage_size = 0.8 args.graph_max_batch_size = 6 logger.info( - f"performance_mode is personal, set running_max_req_size to 3," - f"batch_max_tokens to 2048, chunked_prefill_size to 1024," - f"graph_max_batch_size to 32" + "performance_mode is personal, set running_max_req_size to 3," + "batch_max_tokens to 2048, chunked_prefill_size to 1024," + "graph_max_batch_size to 32" ) if not args.disable_shm_warning: @@ -162,13 +191,20 @@ def _launch_subprocesses(args: StartArgs): ) # mtp params check - if args.mtp_mode is not None: + spec_config = SpeculativeConfig.from_args(args) + spec_config.validate() + if spec_config.enabled: if args.mtp_draft_model_dir is None: - args.mtp_draft_model_dir = [args.model_dir] * args.mtp_step - assert args.mtp_step > 0 + assert not spec_config.uses_block_draft_model, ( + f"--mtp_draft_model_dir is required for {spec_config.mode} mode" + ) + args.mtp_draft_model_dir = [args.model_dir] * spec_config.draft_model_count + elif isinstance(args.mtp_draft_model_dir, str): + args.mtp_draft_model_dir = [args.mtp_draft_model_dir] + assert len(args.mtp_draft_model_dir) >= spec_config.draft_model_count + spec_config = normalize_block_mtp_step_from_first_draft_config(args, spec_config) else: assert args.mtp_draft_model_dir is None - assert args.mtp_step == 0 # automatically set visual_dp based on visual_tp and tp. # In visual proxy mode keep the caller-provided visual_dp / visual_tp. diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index b268120c90..b0a9b59563 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -1,4 +1,3 @@ -import os import math import ctypes import asyncio @@ -133,6 +132,8 @@ class Req(ctypes.Structure): ("cumlogprob", ctypes.c_float), # mtp draft model 多输出命中接受的token数量 ("mtp_accepted_token_num", ctypes.c_int), + ("mtp_verify_token_num", ctypes.c_int), + ("mtp_verify_step_num", ctypes.c_int), # mtp_step 保存一个mtp使用的常量参数,用于快速访问,不会被外部输入初始化 ("_mtp_step", ctypes.c_int), # stop_str_matched 用于判断停止字符串是否匹配成功, detokenization 进程写入,router 进程读取 @@ -203,6 +204,8 @@ def init( self.chunked_prefill_size = chunked_prefill_size self.shm_prompt_ids.arr[0 : len(prompt_ids)] = prompt_ids self.mtp_accepted_token_num = 0 + self.mtp_verify_token_num = 0 + self.mtp_verify_step_num = 0 self._mtp_step = get_env_start_args().mtp_step self.stop_str_matched = False self.stop_str_matched_token_index = -1 diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 758099a05d..d5c5a8dd06 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -152,6 +152,9 @@ class StartArgs: graph_max_batch_size: int = field(default=256) graph_split_batch_size: int = field(default=32) graph_grow_step_size: int = field(default=16) + mtp_draft_graph_max_batch_size: Optional[int] = field(default=None) + mtp_draft_graph_split_batch_size: Optional[int] = field(default=None) + mtp_draft_graph_grow_step_size: Optional[int] = field(default=None) graph_max_len_in_batch: int = field(default=0) quant_type: Optional[str] = field(default="none") quant_cfg: Optional[str] = field(default=None) @@ -187,12 +190,18 @@ class StartArgs: "eagle_with_att", "vanilla_no_att", "eagle_no_att", + "eagle3", + "dspark", + "dflash", + "qwen3next_vanilla", + "qwen3next_eagle", None, ] }, ) - mtp_draft_model_dir: Optional[str] = field(default=None) + mtp_draft_model_dir: Optional[List[str]] = field(default=None) mtp_step: int = field(default=0) + mtp_dynamic_verify: bool = field(default=False) kv_quant_calibration_config_path: Optional[str] = field(default=None) pd_kv_page_num: int = field(default=16) pd_kv_page_size: int = field(default=1024) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 1de9486195..94dccaafd0 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -27,7 +27,7 @@ from lightllm.server.core.objs.out_token_circlequeue import LIGHTLLM_OUT_TOKEN_QUEUE_SIZE from lightllm.server.core.objs.io_objs import GroupReqObjs from lightllm.server.core.objs.shm_req_manager import ShmReqManager -from lightllm.server.core.objs.atomic_array_lock import AtomicShmArrayLock, AsyncLock, AtomicLockItem +from lightllm.server.core.objs.atomic_array_lock import AtomicShmArrayLock, AsyncLock from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt from lightllm.utils.log_utils import init_logger from lightllm.server.metrics.manager import MetricClient @@ -38,6 +38,7 @@ from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.shm_port_args import get_shm_port_args from lightllm.utils.error_utils import ClientDisconnected, PDPrefillNodeStopGenToken +from lightllm.common.speculative import SpeculativeConfig from rpyc.utils.classic import obtain logger = init_logger(__name__) @@ -708,6 +709,11 @@ async def _wait_to_token_package( unfinished_count = sampling_params.best_of out_token_counter = 0 sub_req_id_to_mtp_accepted_token_num: Dict[int, int] = {} + sub_req_id_to_mtp_verify_token_num: Dict[int, int] = {} + sub_req_id_to_mtp_verify_step_num: Dict[int, int] = {} + mtp_verify_step_num = 0 + spec_config = SpeculativeConfig.from_args(self.args) + is_static_mtp = spec_config.step > 0 and not spec_config.dynamic_verify first_token_cost_ms = sys.float_info.max prompt_tokens = len(prompt_ids) is_first_token = True @@ -741,6 +747,15 @@ async def _wait_to_token_package( disk_prompt_cache_len = metadata.pop("disk_prompt_cache_len", 0) metadata["prompt_cache_len"] = gpu_prompt_cache_len + cpu_prompt_cache_len + disk_prompt_cache_len sub_req_id_to_mtp_accepted_token_num[sub_req_id] = metadata.get("mtp_accepted_token_num", 0) + cur_mtp_verify_token_num = metadata.get("mtp_verify_token_num", 0) + prev_mtp_verify_token_num = sub_req_id_to_mtp_verify_token_num.get(sub_req_id, 0) + sub_req_id_to_mtp_verify_token_num[sub_req_id] = cur_mtp_verify_token_num + cur_mtp_verify_step_num = metadata.get("mtp_verify_step_num", 0) + sub_req_id_to_mtp_verify_step_num[sub_req_id] = cur_mtp_verify_step_num + if is_static_mtp and cur_mtp_verify_token_num > prev_mtp_verify_token_num: + mtp_verify_step_num += (cur_mtp_verify_token_num - prev_mtp_verify_token_num) // ( + self.args.mtp_step + 1 + ) if is_first_token: first_token_cost_ms = (time.time() - start_time) * 1000 @@ -758,7 +773,8 @@ async def _wait_to_token_package( unfinished_count -= 1 if unfinished_count == 0: - total_cost_time_ms = (time.time() - start_time) * 1000 + finish_time = time.time() + total_cost_time_ms = (finish_time - start_time) * 1000 mean_per_token_cost_time_ms = (total_cost_time_ms - first_token_cost_ms) / out_token_counter self.per_token_costs.add(mean_per_token_cost_time_ms) x_request_id = request.headers.get("X-Request-Id", "") if request is not None else "" @@ -770,15 +786,32 @@ async def _wait_to_token_package( prompt_cache_ratio = prompt_cache_len / prompt_tokens generation_throughput = out_token_counter / max(total_cost_time_ms / 1000.0, 1e-6) - mtp_avg_token_per_step = out_token_counter / max( - (out_token_counter - sum(sub_req_id_to_mtp_accepted_token_num.values())), 1 - ) + mtp_accepted_token_num = sum(sub_req_id_to_mtp_accepted_token_num.values()) + mtp_verify_token_num = sum(sub_req_id_to_mtp_verify_token_num.values()) + mtp_total_step = sum(sub_req_id_to_mtp_verify_step_num.values()) + if mtp_total_step <= 0: + mtp_total_step = out_token_counter - mtp_accepted_token_num + if mtp_total_step <= 0 and is_static_mtp and mtp_verify_step_num > 0: + mtp_total_step = mtp_verify_step_num + mtp_avg_token_per_step = out_token_counter / max(mtp_total_step, 1) + mtp_avg_verify_tokens_per_step = mtp_verify_token_num / max(mtp_total_step, 1) + mtp_avg_accept_len_per_step_direct = mtp_accepted_token_num / max(mtp_total_step, 1) + decode_start_time = start_time + first_token_cost_ms / 1000.0 + decode_end_time = finish_time + decode_total_time_ms = max((decode_end_time - decode_start_time) * 1000, 0.0) + decode_token_counter = max(out_token_counter - 1, 0) + decode_token_throughput = decode_token_counter / max(decode_total_time_ms / 1000.0, 1e-6) format_start_time = datetime.datetime.fromtimestamp(start_time).strftime("%Y-%m-%d %H:%M:%S") logger.info( f"X-Request-Id:{x_request_id} " f"X-Session-Id:{x_session_id} start_time:{format_start_time} " f"lightllm_req_id:{group_request_id} first_token_cost:{first_token_cost_ms}ms " f"total_cost_time:{total_cost_time_ms}ms,out_token_counter:{out_token_counter} " + f"decode_start_time:{decode_start_time} " + f"decode_end_time:{decode_end_time} " + f"decode_total_time:{decode_total_time_ms}ms " + f"decode_token_counter:{decode_token_counter} " + f"decode_token_throughput:{decode_token_throughput} " f"mean_per_token_cost_time: {mean_per_token_cost_time_ms}ms " f"prompt_token_num:{prompt_tokens} " f"gpu cache hit: {gpu_prompt_cache_ratio > 0} " @@ -790,7 +823,12 @@ async def _wait_to_token_package( f"disk cache hit: {disk_prompt_cache_len > 0} " f"disk_prompt_cache_len:{disk_prompt_cache_len} " f"disk_prompt_cache_ratio:{disk_prompt_cache_ratio} " + f"mtp_accepted_token_num:{mtp_accepted_token_num} " + f"mtp_total_step:{mtp_total_step} " + f"mtp_total_verify_tokens:{mtp_verify_token_num} " f"mtp_avg_token_per_step:{mtp_avg_token_per_step} " + f"mtp_avg_accept_len_per_step_direct:{mtp_avg_accept_len_per_step_direct} " + f"mtp_avg_verify_tokens_per_step:{mtp_avg_verify_tokens_per_step} " ) self.metric_client.histogram_observe("lightllm_cache_length", prompt_cache_len) @@ -935,6 +973,8 @@ async def handle_loop(self): "cpu_prompt_cache_len": req.cpu_prompt_cache_len, "disk_prompt_cache_len": req.disk_prompt_cache_len, "mtp_accepted_token_num": req.mtp_accepted_token_num, + "mtp_verify_token_num": req.mtp_verify_token_num, + "mtp_verify_step_num": req.mtp_verify_step_num, } metadata["logprobs"] = req.get_output_logprobs_metadata(src_index, self.tokenizer) if self.args.use_reward_model: diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 99f123e665..c3ee798f06 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -4,7 +4,6 @@ import uvloop import time import datetime -import ujson as json import pickle import httpx from contextlib import aclosing @@ -391,6 +390,7 @@ async def _wait_to_token_package( unfinished_count = sampling_params.best_of is_first_token = True sub_req_id_to_mtp_accepted_token_num: Dict[int, int] = {} + sub_req_id_to_mtp_verify_step_num: Dict[int, int] = {} async for sub_req_id, out_str, metadata, finish_status in self.fetch_pd_stream( p_node, d_node, prompt, sampling_params, multimodal_params, request @@ -404,6 +404,7 @@ async def _wait_to_token_package( out_token_counter += 1 prompt_cache_len = max(prompt_cache_len, metadata.get("prompt_cache_len", 0)) sub_req_id_to_mtp_accepted_token_num[sub_req_id] = metadata.get("mtp_accepted_token_num", 0) + sub_req_id_to_mtp_verify_step_num[sub_req_id] = metadata.get("mtp_verify_step_num", 0) if is_first_token: first_token_cost_ms = (time.time() - start_time) * 1000 is_first_token = False @@ -422,9 +423,10 @@ async def _wait_to_token_package( x_request_id = request.headers.get("X-Request-Id", "") x_session_id = request.headers.get("X-Session-Id", "") prompt_cache_ratio = prompt_cache_len / prompt_tokens - mtp_avg_token_per_step = out_token_counter / max( - (out_token_counter - sum(sub_req_id_to_mtp_accepted_token_num.values())), 1 - ) + mtp_total_step = sum(sub_req_id_to_mtp_verify_step_num.values()) + if mtp_total_step <= 0: + mtp_total_step = out_token_counter - sum(sub_req_id_to_mtp_accepted_token_num.values()) + mtp_avg_token_per_step = out_token_counter / max(mtp_total_step, 1) format_start_time = datetime.datetime.fromtimestamp(start_time).strftime("%Y-%m-%d %H:%M:%S") logger.info( f"X-Request-Id:{x_request_id} " diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index e0a7ebae77..8f9273336c 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -1,23 +1,20 @@ import enum import torch import torch.distributed as dist -import numpy as np import collections import pickle from sortedcontainers import SortedDict -from dataclasses import dataclass, field +from dataclasses import dataclass from typing import List, Dict, Tuple, Optional, Callable, Any, Union from lightllm.common.req_manager import ReqManager, ReqManagerForMamba -from lightllm.utils.infer_utils import mark_start, mark_end -from lightllm.server.core.objs import Req, SamplingParams, FinishStatus, ShmReqManager +from lightllm.server.core.objs import Req, FinishStatus, ShmReqManager from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache, TreeNode from lightllm.server.router.dynamic_prompt.linear_att_radix_cache import ( LinearAttPagedRadixCache, LinearAttPagedTreeNode, ) from lightllm.utils.log_utils import init_logger -from lightllm.server.req_id_generator import convert_sub_id_to_group_id from lightllm.server.multimodal_params import MultimodalParams from lightllm.utils.custom_kernel_utis import custom_cat from lightllm.utils.envs_utils import get_env_start_args @@ -37,6 +34,7 @@ class InferenceContext: infer_req_ids = None vocab_size = None cpu_embed_cache_client: Optional[CpuEmbedCacheClient] = None + dynamic_mtp_planner: Optional[Any] = None overlap_stream: torch.cuda.Stream = None # 一些情况下推理进程进行异步折叠操作的异步流对象。 cpu_kv_cache_stream: torch.cuda.Stream = None # 用 cpu kv cache 操作的 stream @@ -72,6 +70,47 @@ def init_cpu_embed_cache_client(self): self.cpu_embed_cache_client = CpuEmbedCacheClient(create_meta_data=False, init_shm_data=False) return + def init_dynamic_mtp_planner(self, mtp_step: int, mode: str = None): + if mode == "dspark": + planner_mode = "dspark" + elif mode == "eagle3": + planner_mode = "eagle3" + else: + planner_mode = "default" + if ( + self.dynamic_mtp_planner is not None + and self.dynamic_mtp_planner.mtp_step == mtp_step + and getattr(self.dynamic_mtp_planner, "planner_mode", "default") == planner_mode + ): + return + + from lightllm.server.router.model_infer.speculative.planner import ( + DSparkDynamicMTPPlanner, + DynamicMTPPlanner, + Eagle3DynamicMTPPlanner, + ) + + planner_cls = { + "default": DynamicMTPPlanner, + "dspark": DSparkDynamicMTPPlanner, + "eagle3": Eagle3DynamicMTPPlanner, + }[planner_mode] + self.dynamic_mtp_planner = planner_cls(mtp_step=mtp_step) + return + + def record_dynamic_mtp_infer_cost(self, *, batch_size: int, infer_cost_ms: float, is_draft_model: bool): + if self.dynamic_mtp_planner is None: + self.init_dynamic_mtp_planner( + mtp_step=get_env_start_args().mtp_step, + mode=get_env_start_args().mtp_mode, + ) + self.dynamic_mtp_planner.update_infer_cost( + batch_size=batch_size, + infer_cost_ms=infer_cost_ms, + is_draft_model=is_draft_model, + ) + return + def get_overlap_stream(self) -> torch.cuda.Stream: if self.overlap_stream is None: self.overlap_stream = torch.cuda.Stream() @@ -569,6 +608,7 @@ def __init__( # mtp_step 用来记录一个请求 draft模型每步需要生成的token数量 # 正常模式下,这个值为0,在 mtp 模式下,这个值为 draft 模型每步需要生成的token数量 self.mtp_step: int = get_env_start_args().mtp_step + if self.mtp_step > 0: self.decode_need_token_num = self._mtp_decode_need_token_num else: @@ -859,6 +899,14 @@ def update_mtp_accepted_token_num(self, accept_token_num: int): # 用于统计 mtp 的接受率 self.shm_req.mtp_accepted_token_num += accept_token_num + def update_mtp_verify_token_num(self, verify_token_num: int): + # 用于统计 mtp 验证时发送给主模型的 token 总数 + self.shm_req.mtp_verify_token_num += verify_token_num + + def update_mtp_verify_step_num(self, verify_step_num: int): + # 用于统计 mtp 验证轮数 + self.shm_req.mtp_verify_step_num += verify_step_num + def get_last_gen_token(self): return self.shm_req.shm_prompt_ids.arr[self.shm_req.input_len + self.cur_output_len - 1] diff --git a/lightllm/server/router/model_infer/mode_backend/__init__.py b/lightllm/server/router/model_infer/mode_backend/__init__.py index 8c608e5f92..f26b329856 100644 --- a/lightllm/server/router/model_infer/mode_backend/__init__.py +++ b/lightllm/server/router/model_infer/mode_backend/__init__.py @@ -1,15 +1,46 @@ -from .chunked_prefill.impl import ChunkedPrefillBackend -from .chunked_prefill.impl_for_first_token_constraint_mode import FirstTokenConstraintBackend -from .chunked_prefill.impl_for_outlines_constraint_mode import OutlinesConstraintBackend -from .chunked_prefill.impl_for_reward_model import RewardModelBackend -from .chunked_prefill.impl_for_token_healing import TokenHealingBackend -from .chunked_prefill.impl_for_xgrammar_mode import XgrammarBackend - -from .dp_backend.impl import DPChunkedPrefillBackend -from .diverse_backend.impl import DiversehBackend - -# pd mode backend -from .pd.prefill_node_impl.prefill_impl import PDChunkedPrefillForPrefillNode -from .pd.prefill_node_impl.prefill_impl_for_dp import PDDPChunkedForPrefillNode -from .pd.decode_node_impl.decode_impl import PDDecodeNode -from .pd.decode_node_impl.decode_impl_for_dp import PDDPForDecodeNode +from importlib import import_module + + +_BACKEND_EXPORTS = { + "ChunkedPrefillBackend": (".chunked_prefill.impl", "ChunkedPrefillBackend"), + "FirstTokenConstraintBackend": ( + ".chunked_prefill.impl_for_first_token_constraint_mode", + "FirstTokenConstraintBackend", + ), + "OutlinesConstraintBackend": ( + ".chunked_prefill.impl_for_outlines_constraint_mode", + "OutlinesConstraintBackend", + ), + "ReturnPromptLogProbBackend": ( + ".chunked_prefill.impl_for_return_all_prompt_logprobs", + "ReturnPromptLogProbBackend", + ), + "RewardModelBackend": (".chunked_prefill.impl_for_reward_model", "RewardModelBackend"), + "TokenHealingBackend": (".chunked_prefill.impl_for_token_healing", "TokenHealingBackend"), + "XgrammarBackend": (".chunked_prefill.impl_for_xgrammar_mode", "XgrammarBackend"), + "DPChunkedPrefillBackend": (".dp_backend.impl", "DPChunkedPrefillBackend"), + "DiversehBackend": (".diverse_backend.impl", "DiversehBackend"), + "PDChunkedPrefillForPrefillNode": ( + ".pd.prefill_node_impl.prefill_impl", + "PDChunkedPrefillForPrefillNode", + ), + "PDDPChunkedForPrefillNode": ( + ".pd.prefill_node_impl.prefill_impl_for_dp", + "PDDPChunkedForPrefillNode", + ), + "PDDecodeNode": (".pd.decode_node_impl.decode_impl", "PDDecodeNode"), + "PDDPForDecodeNode": (".pd.decode_node_impl.decode_impl_for_dp", "PDDPForDecodeNode"), +} + + +def __getattr__(name): + if name not in _BACKEND_EXPORTS: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + module_name, attr_name = _BACKEND_EXPORTS[name] + value = getattr(import_module(module_name, __name__), attr_name) + globals()[name] = value + return value + + +__all__ = list(_BACKEND_EXPORTS) 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 eaf3552607..d9e9b160ed 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -4,7 +4,9 @@ import time import threading import torch.distributed as dist -from typing import List, Tuple, Callable, Optional, Union +import collections +from dataclasses import replace +from typing import List, Tuple, Callable, Optional from transformers.configuration_utils import PretrainedConfig from lightllm.utils.infer_utils import set_random_seed from lightllm.utils.log_utils import init_logger @@ -19,7 +21,6 @@ from lightllm.server.router.dynamic_prompt.linear_att_radix_cache import LinearAttPagedRadixCache from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache from lightllm.common.basemodel.batch_objs import ModelOutput, ModelInput -from lightllm.common.basemodel.triton_kernel.mtp_utils import mtp_verify from lightllm.utils.dist_utils import init_distributed_env from lightllm.utils.envs_utils import get_unique_server_name from lightllm.server.core.objs import ShmReqManager, StartArgs @@ -34,8 +35,19 @@ get_env_start_args, enable_radix_tree_timer_merge, get_radix_tree_merge_update_delta, + enable_dynamic_mtp_verify, ) from lightllm.distributed import dist_group_manager +from lightllm.common.speculative import ( + SpeculativeConfig, + get_dspark_family_block_size, + is_dspark_draft_config, + is_eagle3_draft_config, + is_gemma4_dspark_draft_config, + is_qwen3_dflash_draft_config, + is_qwen3_dspark_draft_config, +) +from lightllm.server.router.model_infer.speculative import build_spec_runtime from lightllm.distributed.communication_op import ( all_gather_into_tensor, all_reduce, @@ -53,10 +65,13 @@ from .multi_level_kv_cache import MultiLevelKvCacheModule from lightllm.utils.profiler import ProcessProfiler, ProfilerCmd +logger = init_logger(__name__) + class ModeBackend: def __init__(self) -> None: self.shm_req_manager = ShmReqManager() + start_args = get_env_start_args() self.overlap_event_manager = OverlapEventManager() # 标识是否支持 overlap 功能,很多子类模式如 xgrammar 和 outlines 当前不支持 overlap 高性能模式 @@ -69,8 +84,11 @@ def __init__(self) -> None: # extra_post_req_handle_func 用于添加请求InferReq的状态变化中添加额外的后处理信息,主要是状态机相关的调整等。 self.extra_post_req_handle_func: Optional[Callable[[InferReq, int, float], None]] = None - self.enable_decode_microbatch_overlap = get_env_start_args().enable_decode_microbatch_overlap - self.enable_prefill_microbatch_overlap = get_env_start_args().enable_prefill_microbatch_overlap + self.enable_decode_microbatch_overlap = start_args.enable_decode_microbatch_overlap + self.enable_prefill_microbatch_overlap = start_args.enable_prefill_microbatch_overlap + self.spec_config = SpeculativeConfig.from_args(start_args, dynamic_verify=enable_dynamic_mtp_verify()) + self.spec_config.validate() + self.spec_adapter = None # 控制 _get_classed_reqs 分类的参数变量,不同的 backend 具有可能需要不同的分类运行条件。 self.classed_req_no_decode = False @@ -85,6 +103,7 @@ def __init__(self) -> None: self._radix_tree_merge_update_delta: int = get_radix_tree_merge_update_delta() pass + def init_model(self, kvargs): self.args: StartArgs = kvargs.get("args", None) assert self.args is not None @@ -108,10 +127,22 @@ def init_model(self, kvargs): self.is_multinode_tp = self.args.nnodes > 1 and self.args.dp == 1 self.is_pd_mode = self.run_mode in ["prefill", "decode"] self.is_pd_decode_mode = self.run_mode == "decode" + self.spec_config = SpeculativeConfig.from_args(self.args, dynamic_verify=enable_dynamic_mtp_verify()) + self.spec_config.validate() + if self.spec_config.needs_target_layer_hidden: + assert ( + not self.args.enable_decode_microbatch_overlap + ), f"{self.spec_config.mode} mode does not support decode microbatch overlap" + assert ( + not self.args.enable_prefill_microbatch_overlap + ), f"{self.spec_config.mode} mode does not support prefill microbatch overlap" self.logger = init_logger(__name__) self.weight_dir = kvargs["weight_dir"] + self._normalize_block_mtp_step_from_first_draft_config() + # p d 分离模式,decode节点才会使用的参数 + self.pd_rpyc_ports = kvargs.get("pd_rpyc_ports", None) max_total_token_num = kvargs["max_total_token_num"] init_distributed_env(kvargs) @@ -139,6 +170,8 @@ def init_model(self, kvargs): "disable_chunked_prefill": self.disable_chunked_prefill, "data_type": kvargs.get("data_type", "float16"), "graph_max_batch_size": kvargs.get("graph_max_batch_size", 16), + "graph_split_batch_size": kvargs.get("graph_split_batch_size", self.args.graph_split_batch_size), + "graph_grow_step_size": kvargs.get("graph_grow_step_size", self.args.graph_grow_step_size), "graph_max_len_in_batch": kvargs.get("graph_max_len_in_batch", 8196), "disable_cudagraph": kvargs.get("disable_cudagraph", False), "mem_fraction": kvargs.get("mem_fraction", 0.9), @@ -246,8 +279,12 @@ def init_model(self, kvargs): self.shm_pd_trans_io_buffer = ShmObjsIOBuffer(tail_str="pd") # 开启 mtp 模式,需要完成mtp model的初始化 - if self.args.mtp_mode: + if self.spec_config.enabled: self.init_mtp_draft_model(kvargs) + if self.spec_config.dynamic_verify: + g_infer_context.init_dynamic_mtp_planner(mtp_step=self.mtp_step, mode=self.spec_config.mode) + self.spec_adapter = build_spec_runtime(self) + self._attach_spec_adapter() if self.args.enable_cpu_cache: self.multi_level_cache_module = MultiLevelKvCacheModule(self) @@ -303,21 +340,43 @@ def decode(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): def init_mtp_draft_model(self, main_kvargs: dict): self.mtp_step = self.args.mtp_step self.draft_models = [] + spec_config = self.spec_config os.environ["DISABLE_CHECK_MAX_LEN_INFER"] = "1" - if self.args.mtp_mode in ["vanilla_with_att", "vanilla_no_att"]: - num_mtp_modules = self.args.mtp_step - elif self.args.mtp_mode in ["eagle_with_att", "eagle_no_att"]: - num_mtp_modules = 1 - else: - assert False, f"error mtp mode {self.args.mtp_mode}" + num_mtp_modules = spec_config.draft_model_count + mtp_draft_model_dirs = self.args.mtp_draft_model_dir + if isinstance(mtp_draft_model_dirs, str): + mtp_draft_model_dirs = [mtp_draft_model_dirs] + assert mtp_draft_model_dirs is not None + assert len(mtp_draft_model_dirs) >= num_mtp_modules + + draft_graph_max_override = getattr(self.args, "mtp_draft_graph_max_batch_size", None) + draft_graph_split_override = getattr(self.args, "mtp_draft_graph_split_batch_size", None) + draft_graph_grow_override = getattr(self.args, "mtp_draft_graph_grow_step_size", None) + draft_graph_max_batch_size = ( + draft_graph_max_override + if draft_graph_max_override is not None + else main_kvargs.get("graph_max_batch_size", 16) + ) + draft_graph_split_batch_size = ( + draft_graph_split_override + if draft_graph_split_override is not None + else main_kvargs.get("graph_split_batch_size", self.args.graph_split_batch_size) + ) + draft_graph_grow_step_size = ( + draft_graph_grow_override + if draft_graph_grow_override is not None + else main_kvargs.get("graph_grow_step_size", self.args.graph_grow_step_size) + ) for i in range(num_mtp_modules): - mtp_model_cfg, _ = PretrainedConfig.get_config_dict(self.args.mtp_draft_model_dir[i]) + mtp_model_cfg, _ = PretrainedConfig.get_config_dict(mtp_draft_model_dirs[i]) + self._normalize_block_mtp_step_from_config(mtp_model_cfg) + spec_config = self.spec_config model_type = mtp_model_cfg.get("model_type", "") mtp_model_kvargs = { - "weight_dir": self.args.mtp_draft_model_dir[i], + "weight_dir": mtp_draft_model_dirs[i], "max_total_token_num": self.model.mem_manager.size, "load_way": main_kvargs["load_way"], "max_req_num": main_kvargs.get("max_req_num", 1000), @@ -326,7 +385,9 @@ def init_mtp_draft_model(self, main_kvargs: dict): "return_all_prompt_logics": False, "disable_chunked_prefill": self.disable_chunked_prefill, "data_type": main_kvargs.get("data_type", "float16"), - "graph_max_batch_size": main_kvargs.get("graph_max_batch_size", 16), + "graph_max_batch_size": draft_graph_max_batch_size, + "graph_split_batch_size": draft_graph_split_batch_size, + "graph_grow_step_size": draft_graph_grow_step_size, "graph_max_len_in_batch": main_kvargs.get("graph_max_len_in_batch", 8196), "disable_cudagraph": main_kvargs.get("disable_cudagraph", False), "mem_fraction": main_kvargs["mem_fraction"], @@ -341,33 +402,82 @@ def init_mtp_draft_model(self, main_kvargs: dict): model_type = mtp_model_cfg.get("model_type", "") if model_type == "deepseek_v3": - assert self.args.mtp_mode in ["vanilla_with_att", "eagle_with_att"] + assert spec_config.uses_attention_draft self.draft_models.append(Deepseek3MTPModel(mtp_model_kvargs)) elif model_type == "qwen3_moe": - assert self.args.mtp_mode in ["vanilla_no_att", "eagle_no_att"] + assert spec_config.uses_no_attention_draft and not spec_config.is_eagle3 self.draft_models.append(Qwen3MOEMTPModel(mtp_model_kvargs)) elif model_type == "mistral": - assert self.args.mtp_mode in ["vanilla_no_att", "eagle_no_att"] + assert spec_config.uses_no_attention_draft and not spec_config.is_eagle3 self.draft_models.append(MistralMTPModel(mtp_model_kvargs)) elif model_type == "glm4_moe_lite": - assert self.args.mtp_mode in ["vanilla_with_att", "eagle_with_att"] + assert spec_config.uses_attention_draft self.draft_models.append(Glm4MoeLiteMTPModel(mtp_model_kvargs)) elif model_type in ("qwen3_5", "qwen3_5_text"): - assert self.args.mtp_mode in ["vanilla_with_att", "eagle_with_att"] + assert spec_config.uses_attention_draft from lightllm.models.qwen3_5_mtp.model import Qwen3_5MTPModel self.draft_models.append(Qwen3_5MTPModel(mtp_model_kvargs)) elif model_type in ("qwen3_5_moe", "qwen3_5_moe_text"): - assert self.args.mtp_mode in ["vanilla_with_att", "eagle_with_att"] + assert spec_config.uses_attention_draft from lightllm.models.qwen3_5_moe_mtp.model import Qwen3_5MoeMTPModel self.draft_models.append(Qwen3_5MoeMTPModel(mtp_model_kvargs)) + elif spec_config.is_eagle3 and is_eagle3_draft_config(mtp_model_cfg): + from lightllm.models.qwen3_eagle.model import Qwen3EagleModel + + self.draft_models.append(Qwen3EagleModel(mtp_model_kvargs)) + elif spec_config.is_dflash and is_qwen3_dflash_draft_config(mtp_model_cfg): + from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel + + self.draft_models.append(Qwen3DFlashModel(mtp_model_kvargs)) + elif spec_config.is_dspark and is_qwen3_dspark_draft_config(mtp_model_cfg): + from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel + + self.draft_models.append(Qwen3DSparkModel(mtp_model_kvargs)) + elif (spec_config.is_dflash or spec_config.is_dspark) and is_gemma4_dspark_draft_config(mtp_model_cfg): + raise NotImplementedError("Gemma4 DSpark draft checkpoints are not wired to LightLLM serving yet.") + elif (spec_config.is_dflash or spec_config.is_dspark) and is_dspark_draft_config(mtp_model_cfg): + raise ValueError(f"Unsupported DSpark-family draft architecture: {mtp_model_cfg.get('architectures')}") else: raise ValueError(f"Unsupported MTP model type: {model_type}") self.logger.info(f"loaded mtp model class {self.draft_models[i].__class__}") return + def _normalize_block_mtp_step_from_config(self, mtp_model_cfg: dict) -> None: + if not self.spec_config.uses_block_draft_model: + return + + block_size = get_dspark_family_block_size( + mtp_model_cfg, + require_confidence_head=self.spec_config.is_dspark, + ) + configured_step = int(getattr(self.args, "mtp_step", 0)) + if configured_step not in (0, block_size): + self.logger.warning( + "Overriding mtp_step=%s with block draft config block_size=%s for %s mode", + configured_step, + block_size, + self.spec_config.mode, + ) + self.args.mtp_step = block_size + self.mtp_step = block_size + self.spec_config = replace(self.spec_config, step=block_size) + return + + def _normalize_block_mtp_step_from_first_draft_config(self) -> None: + if not self.spec_config.uses_block_draft_model: + return + + mtp_draft_model_dirs = self.args.mtp_draft_model_dir + if isinstance(mtp_draft_model_dirs, str): + mtp_draft_model_dirs = [mtp_draft_model_dirs] + assert mtp_draft_model_dirs is not None and len(mtp_draft_model_dirs) > 0 + mtp_model_cfg, _ = PretrainedConfig.get_config_dict(mtp_draft_model_dirs[0]) + self._normalize_block_mtp_step_from_config(mtp_model_cfg) + return + def _async_copy_next_token_infos_to_pin_mem( self, next_token_ids: torch.Tensor, @@ -470,6 +580,15 @@ def _capture_prompt_logprobs_if_needed( start_loc += q_len return + def _attach_spec_adapter(self) -> None: + if not self.spec_config.enabled: + return + assert self.spec_adapter is not None + self.model.set_spec_adapter(self.spec_adapter) + for draft_model in self.draft_models: + draft_model.set_spec_adapter(self.spec_adapter) + return + def _try_read_new_reqs(self): if self.is_multinode_tp: self._try_read_new_reqs_multinode_tp() @@ -823,16 +942,30 @@ def _post_handle( run_reqs_update_packs: List[InferReqUpdatePack], extra_post_req_handle_func: Optional[Callable[[InferReq, int, float], None]] = None, pd_prefill_chunked_handle_func: Optional[Callable[[InferReq, int, float, int], None]] = None, + count_mtp_accepted_tokens: bool = False, ): """ extra_post_req_handle_func 用于提供在一个请求确定输出的时候,给出额外的后处理操作,主要是用于 约束输出等模式,设置自己请求内部的状态机的状态,并添加额外的停止判定条件等。 """ + if isinstance(next_token_ids, torch.Tensor): + next_token_ids = next_token_ids.numpy() + if isinstance(next_token_logprobs, torch.Tensor): + next_token_logprobs = next_token_logprobs.numpy() + if isinstance(next_token_ranks, torch.Tensor): + next_token_ranks = next_token_ranks.numpy() + + mtp_seen_count = collections.Counter() if count_mtp_accepted_tokens and self.is_master_in_dp else None for req_obj, next_token_id, next_token_logprob, next_token_rank, pack in zip( run_reqs, next_token_ids, next_token_logprobs, next_token_ranks, run_reqs_update_packs ): req_obj: InferReq = req_obj pack: InferReqUpdatePack = pack + if mtp_seen_count is not None: + seen_count = mtp_seen_count[req_obj.req_idx] + mtp_seen_count[req_obj.req_idx] += 1 + if seen_count > 0 and not req_obj.finish_status.is_finished(): + req_obj.update_mtp_accepted_token_num(accept_token_num=1) pack.handle( next_token_id=next_token_id, next_token_logprob=next_token_logprob, @@ -858,32 +991,56 @@ def _filter_reqs(self, reqs: List[InferReq]): def _trans_req_ids_to_req_objs(self, req_ids: List[int]) -> List[InferReq]: return [g_infer_context.requests_mapping[req_id] for req_id in req_ids] - def _verify_mtp_v2( - self, new_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, b_req_mtp_start_loc: torch.Tensor - ): - mtp_accept_len, accepted_index = mtp_verify( - req_to_next_token_ids=self.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - new_next_token_ids=new_next_token_ids, - b_req_idx=b_req_idx, - ) - return mtp_accept_len, accepted_index - def _update_mtp_accept_ratio( self, decode_reqs: List[InferReq], mtp_accept_len_cpu: torch.Tensor, ): if self.is_master_in_dp: - for req, accept_len in zip(decode_reqs, mtp_accept_len_cpu): - req.update_mtp_accepted_token_num(accept_token_num=accept_len - 1) + for req, accept_len in zip(decode_reqs, mtp_accept_len_cpu.numpy()): + req.update_mtp_accepted_token_num(accept_token_num=max(int(accept_len) - 1, 0)) + + return + + def _update_mtp_verify_token_num( + self, decode_reqs: List[InferReq], dynamic_mtp_run_reqs: Optional[List[InferReq]] = None + ): + if self.is_master_in_dp: + if dynamic_mtp_run_reqs is None: + for req in decode_reqs: + assert req.mtp_step > 0 + verify_len = 1 + req.mtp_step + req.update_mtp_verify_token_num(verify_token_num=verify_len) + req.update_mtp_verify_step_num(verify_step_num=1) + else: + counter = collections.Counter([req.req_idx for req in dynamic_mtp_run_reqs]) + for req in decode_reqs: + verify_token_num = counter[req.req_idx] + if verify_token_num > 0: + req.update_mtp_verify_token_num(verify_token_num=verify_token_num) + req.update_mtp_verify_step_num(verify_step_num=1) return def _gen_argmax_token_ids(self, model_output: ModelOutput): logits = model_output.logits draft_next_token_ids_gpu = torch.argmax(logits, dim=-1) + + # 如果draft和target的词表不同,需要把draft token映射回主模型词表。 + if self.spec_config.needs_draft_vocab_mapping: + draft_next_token_ids_gpu = self.draft_models[0].map_draft_vocab_to_main_vocab(draft_next_token_ids_gpu) return draft_next_token_ids_gpu + def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): + logits = model_output.logits + probs = torch.softmax(logits, dim=-1) + max_probs, draft_next_token_ids_gpu = torch.max(probs, dim=-1) + + # 如果self.d2t不为None,那么draft的token需要进行相应的转换 + if self.spec_config.needs_draft_vocab_mapping: + draft_next_token_ids_gpu = self.draft_models[0].map_draft_vocab_to_main_vocab(draft_next_token_ids_gpu) + + return draft_next_token_ids_gpu, max_probs + def _sample_and_scatter_token( self, logits: torch.Tensor, diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index d75302800b..c219d55ee9 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -1,7 +1,7 @@ import torch import time -from typing import List, Optional, Callable, Dict, Any -from queue import Queue +import torch.distributed as dist +from typing import List from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.mode_backend.overlap_events import OverlapEventPack from lightllm.server.router.model_infer.infer_batch import InferReq @@ -9,22 +9,12 @@ prepare_prefill_inputs, prepare_decode_inputs, ) -from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import ( - prepare_mtp_prefill_inputs, -) from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample from lightllm.server.router.model_infer.infer_batch import g_infer_context -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -from lightllm.common.basemodel.batch_objs import ModelOutput, ModelInput -from lightllm.common.basemodel.triton_kernel.gather_token_id import scatter_token -from lightllm.common.basemodel.triton_kernel.mtp_utils import ( - linear_att_mtp_state_index_update, - mtp_scatter_next_token_ids, -) from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_current_device_id -from lightllm.utils.envs_utils import get_env_start_args from .control_state import ControlState +from lightllm.utils.dist_utils import create_new_group_for_current_dp logger = init_logger(__name__) @@ -35,14 +25,13 @@ def __init__(self) -> None: # 用于控制每一步是执行prefill 和 decode 还是跳过 self.control_state_machine = ControlState() + self.enable_dynamic_mtp = False # 在 mtp 模式下切换绑定的prefill 和 decode 函数 - if get_env_start_args().mtp_mode: + if self.spec_config.enabled: self.prefill = self.prefill_mtp self.decode = self.decode_mtp - self.is_mtp_eagle = get_env_start_args().mtp_mode in ["eagle_with_att", "eagle_no_att"] - self.num_mtp_models = 1 if self.is_mtp_eagle else get_env_start_args().mtp_step - self._draft_decode_func = self._draft_decode_eagle if self.is_mtp_eagle else self._draft_decode_vanilla + self.enable_dynamic_mtp = self.spec_config.dynamic_verify else: self.prefill = self.prefill_normal self.decode = self.decode_normal @@ -50,6 +39,13 @@ def __init__(self) -> None: self.classed_req_strict_prefill = False return + def init_custom(self): + super().init_custom() + if self.enable_dynamic_mtp: + self.mtp_gloo_group = create_new_group_for_current_dp("gloo") + logger.info(f"mtp_gloo_group ranks {dist.get_rank(self.mtp_gloo_group)}") + return + def infer_loop(self): torch.cuda.set_device(get_current_device_id()) try: @@ -209,8 +205,10 @@ def prefill_mtp( mask_func=self.prefill_mask_func, ) # mtp kv fill - self._draft_prefill_forward( - model_input=model_input, model_output=model_output, next_token_ids=next_token_ids + spec_runtime = self.spec_adapter + spec_runtime.build_initial_draft_state( + model_input=model_input, + next_token_ids=next_token_ids, ) g_infer_context.copy_linear_att_state_to_cache_buffer( b_req_idx=model_input.b_req_idx, @@ -250,197 +248,79 @@ def decode_mtp( MTP解码的通用流程,整合eagle和vanilla的共同逻辑 """ model_input, run_reqs = prepare_decode_inputs(decode_reqs) + spec_runtime = self.spec_adapter with torch.cuda.stream(g_infer_context.get_overlap_stream()): - b_mtp_index_cpu = model_input.b_mtp_index - model_output = self.model.forward(model_input) - next_token_ids, next_token_logprobs = sample(model_output.logits, run_reqs, self.eos_id) - next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids) - # verify the next_token_ids - b_req_mtp_start_loc = [index for index, mtp_index in enumerate(b_mtp_index_cpu) if mtp_index == 0] - b_req_mtp_start_loc = g_pin_mem_manager.gen_from_list( - key="b_req_mtp_start_loc", - data=b_req_mtp_start_loc, - dtype=torch.int32, - ).cuda(non_blocking=True) - - mtp_accept_len, accepted_index = self._verify_mtp_v2( - new_next_token_ids=next_token_ids, - b_req_idx=model_input.b_req_idx, - b_req_mtp_start_loc=b_req_mtp_start_loc, - ) - if self.is_linear_att_mixed_model: - linear_att_mtp_state_index_update( - req_to_mtp_state_index=self.model.req_manager.req_to_mtp_state_index, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=model_input.b_req_idx, - b_mtp_index=model_input.b_mtp_index, - accepted_index=accepted_index, - max_mtp_step=self.mtp_step + 1, - ) - accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="accepted_index", - gpu_tensor=accepted_index, - ) - mtp_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="mtp_accept_len", - gpu_tensor=mtp_accept_len, + spec_plan = spec_runtime.plan_decode(model_input=model_input, req_num=len(decode_reqs)) + + model_input, selected_run_reqs = spec_runtime.prepare_decode_model_input( + model_input=model_input, + req_num=len(decode_reqs), + plan=spec_plan, ) - verify_event = torch.cuda.Event() - verify_event.record() + selected_run_reqs_cpu = spec_runtime.async_copy_selected_run_reqs(selected_run_reqs) - ( - next_token_ids_cpu, - next_token_logprobs_cpu, - next_token_ranks_cpu, - ) = self._async_copy_next_token_infos_to_pin_mem(next_token_ids, next_token_logprobs, next_token_ranks) + model_output = self.model.forward(model_input) - # 调用具体的draft decode函数 - additional_mem_indexes_cpu = self._draft_decode_func( - main_model_input=model_input, - main_model_output=model_output, - next_token_ids=next_token_ids, - mtp_accept_len=mtp_accept_len, - b_req_mtp_start_loc=b_req_mtp_start_loc, + next_token_ids, next_token_logprobs = sample( + model_output.logits, + run_reqs, + self.eos_id, + dynamic_batch_size=spec_plan.dynamic_batch_size, + selected_run_reqs=selected_run_reqs, ) + next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids) - g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( - b_req_idx=model_input.b_req_idx, + spec_decode_state = spec_runtime.run_decode_speculative_forward( + model_input=model_input, + model_output=model_output, + run_reqs=run_reqs, + req_num=len(decode_reqs), + plan=spec_plan, + selected_run_reqs_cpu=selected_run_reqs_cpu, next_token_ids=next_token_ids, - mask=accepted_index == 1, + next_token_logprobs=next_token_logprobs, + next_token_ranks=next_token_ranks, + copy_next_token_infos=self._async_copy_next_token_infos_to_pin_mem, ) - sync_event = torch.cuda.Event() - sync_event.record() # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() - verify_event.synchronize() - verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] + + run_reqs, verify_ok_reqs = spec_runtime.resolve_decode_pre_post_reqs( + state=spec_decode_state, + decode_reqs=decode_reqs, + ) + self._update_mtp_verify_token_num( + decode_reqs=decode_reqs, + dynamic_mtp_run_reqs=run_reqs if self.enable_dynamic_mtp else None, + ) update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) # 第三阶段 event_pack.notify_forward_and_wait_post_handle() - sync_event.synchronize() - - # 处理需要释放的内存索引 - need_free_mem_indexes = model_input.mem_indexes_cpu[accepted_index_cpu == 0] - if additional_mem_indexes_cpu is not None: - need_free_mem_indexes = torch.cat([need_free_mem_indexes, additional_mem_indexes_cpu], dim=0) + spec_post_state = spec_runtime.finish_decode_post( + state=spec_decode_state, + req_num=len(decode_reqs), + run_reqs=run_reqs, + ) + self._update_mtp_accept_ratio( + decode_reqs=decode_reqs, + mtp_accept_len_cpu=spec_post_state.mtp_accept_len_cpu, + ) - self._update_mtp_accept_ratio(decode_reqs=decode_reqs, mtp_accept_len_cpu=mtp_accept_len_cpu) - select_mask = torch.tensor(accepted_index_cpu, dtype=torch.bool, device="cpu") self._post_handle( run_reqs=verify_ok_reqs, - next_token_ids=next_token_ids_cpu[select_mask], - next_token_logprobs=next_token_logprobs_cpu[select_mask], - next_token_ranks=next_token_ranks_cpu[select_mask], + next_token_ids=spec_post_state.next_token_ids, + next_token_logprobs=spec_post_state.next_token_logprobs, + next_token_ranks=spec_post_state.next_token_ranks, run_reqs_update_packs=update_packs, extra_post_req_handle_func=self.extra_post_req_handle_func, ) - if len(need_free_mem_indexes) > 0: - g_infer_context.req_manager.mem_manager.free(need_free_mem_indexes) + if len(spec_post_state.need_free_mem_indexes) > 0: + g_infer_context.req_manager.mem_manager.free(spec_post_state.need_free_mem_indexes) # 第四阶段 event_pack.notify_pre_post_handle() return - - def _draft_prefill_forward(self, model_input: ModelInput, model_output: ModelOutput, next_token_ids: torch.Tensor): - # spec prefill: MTP, 这个地方只是为了填充draft model的 kv, 并不会使用生成的token_id。 - draft_model_input = model_input - draft_model_output = model_output - draft_next_token_ids_gpu = next_token_ids - for draft_model_idx in range(self.num_mtp_models): - draft_model_input = prepare_mtp_prefill_inputs( - model_input=draft_model_input, - b_next_token_ids=draft_next_token_ids_gpu, - mtp_draft_input_hiddens=draft_model_output.mtp_main_output_hiddens, - ) - draft_model_output = self.draft_models[draft_model_idx].forward(draft_model_input) - draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) - return - - def _draft_decode_vanilla( - self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - mtp_accept_len: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - ): - # share some inference info with the main model - draft_model_input = main_model_input - draft_model_output = main_model_output - draft_next_token_ids = next_token_ids - all_next_token_ids = [] - all_next_token_ids.append(next_token_ids) - # process the draft model output - for draft_model_idx in range(self.mtp_step): - - draft_model_input.input_ids = draft_next_token_ids - draft_model_input.mtp_draft_input_hiddens = draft_model_output.mtp_main_output_hiddens - # spec decode: MTP - draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) - draft_next_token_ids = self._gen_argmax_token_ids(draft_model_output) - all_next_token_ids.append(draft_next_token_ids) - - all_next_token_ids = torch.stack(all_next_token_ids, dim=1) # [batch_size, mtp_step + 1] - - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=main_model_input.b_req_idx, - mtp_accept_len=mtp_accept_len, - ) - return None - - def _draft_decode_eagle( - self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - mtp_accept_len: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - ): - batch_size = main_model_input.batch_size - num_reqs = batch_size // (self.mtp_step + 1) - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(num_reqs * self.mtp_step) - eagle_mem_indexes_cpu = g_infer_context.req_manager.mem_manager.alloc(num_reqs * self.mtp_step) - eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) - - # share some inference info with the main model - draft_model_input = main_model_input - draft_model_output = main_model_output - draft_next_token_ids = next_token_ids - all_next_token_ids = [] - all_next_token_ids.append(next_token_ids) - # process the draft model output - for _step in range(self.mtp_step): - - draft_model_input.input_ids = draft_next_token_ids - draft_model_input.mtp_draft_input_hiddens = draft_model_output.mtp_main_output_hiddens - # spec decode: MTP - draft_model_idx = _step % self.num_mtp_models - draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) - draft_next_token_ids = self._gen_argmax_token_ids(draft_model_output) - draft_model_input.b_seq_len += 1 - draft_model_input.max_kv_seq_len += 1 - eagle_mem_indexes_i = eagle_mem_indexes[_step * num_reqs : (_step + 1) * num_reqs] - draft_model_input.mem_indexes = torch.cat( - [draft_model_input.mem_indexes.view(-1, self.mtp_step + 1)[:, 1:], eagle_mem_indexes_i.view(-1, 1)], - dim=1, - ).view(-1) - all_next_token_ids.append(draft_next_token_ids) - - all_next_token_ids = torch.stack(all_next_token_ids, dim=1) # [batch_size, mtp_step + 1] - - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=main_model_input.b_req_idx, - mtp_accept_len=mtp_accept_len, - ) - return eagle_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py index 1edbd30306..5bec2db491 100644 --- a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py @@ -1,5 +1,4 @@ import torch -from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.infer_batch import ( g_infer_context, InferReq, @@ -14,17 +13,16 @@ from lightllm.common.basemodel.triton_kernel.gather_token_id import scatter_token from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from ..chunked_prefill.impl import ChunkedPrefillBackend -from lightllm.utils.envs_utils import get_env_start_args class DiversehBackend(ChunkedPrefillBackend): def __init__(self) -> None: super().__init__() - if get_env_start_args().mtp_mode: + if self.spec_config.enabled: # 当前只有 mistral mtp 可以使用 diverse mode 的 mtp 功能。 self.prefill = self.beam_prefill - assert get_env_start_args().mtp_mode in ["vanilla_no_att", "eagle_no_att"] + assert self.spec_config.uses_no_attention_draft and not self.spec_config.is_eagle3 else: self.prefill = self.beam_prefill diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index f6ca89e651..4ab9c7e7d0 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -1,11 +1,9 @@ import torch import time -import torch.nn.functional as F -import torch.distributed as dist -from typing import List, Tuple, Optional, Callable +from typing import List, Tuple from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.common.basemodel.batch_objs import ModelOutput, ModelInput -from lightllm.server.router.model_infer.infer_batch import InferSamplingParams, g_infer_context, InferReq +from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample from lightllm.server.router.model_infer.mode_backend.pre import ( padded_prepare_prefill_inputs, @@ -14,15 +12,10 @@ padded_overlap_prepare_decode_inputs, ) from lightllm.server.router.model_infer.mode_backend.overlap_events import OverlapEventPack -from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import ( - prepare_mtp_prefill_inputs, -) from lightllm.utils.dist_utils import get_current_device_id -from lightllm.utils.envs_utils import get_env_start_args from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.common.basemodel.triton_kernel.mtp_utils import ( linear_att_mtp_state_index_update, - mtp_scatter_next_token_ids, ) from .control_state import DPControlState @@ -35,9 +28,11 @@ def __init__(self) -> None: self.control_state_machine = DPControlState(backend=self) # 在 mtp 模式下切换绑定的prefill 和 decode 函数 - if get_env_start_args().mtp_mode: - self.is_mtp_eagle = get_env_start_args().mtp_mode in ["eagle_with_att", "eagle_no_att"] - self.num_mtp_models = 1 if self.is_mtp_eagle else get_env_start_args().mtp_step + if self.spec_config.enabled: + if self.spec_config.uses_block_draft_model: + raise NotImplementedError("DP backend does not support DFlash/DSpark block draft mode yet.") + self.is_mtp_eagle = self.spec_config.uses_recurrent_draft_model + self.num_mtp_models = self.spec_config.draft_model_count if self.enable_prefill_microbatch_overlap: self.prefill = self.prefill_overlap_mtp else: @@ -431,12 +426,14 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] ) # mtp kv fill - draft_next_token_ids_gpu = torch.zeros((model_input.batch_size), dtype=torch.int64, device="cuda") - if req_num > 0: - draft_next_token_ids_gpu[0:req_num].copy_(next_token_ids) - self._draft_prefill_forward( + draft_next_token_ids_gpu = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num > 0 else None, + batch_size=model_input.batch_size, + copy_len=req_num, + device=model_input.b_req_idx.device, + ) + self.spec_adapter.build_initial_draft_state( model_input=model_input, - model_output=model_output, next_token_ids=draft_next_token_ids_gpu, ) if req_num > 0: @@ -502,11 +499,13 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): dtype=torch.int32, ).cuda(non_blocking=True) - mtp_accept_len, accepted_index = self._verify_mtp_v2( + verify_result = self.spec_adapter.verify_target_tokens( new_next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, ) + mtp_accept_len = verify_result.accept_len + accepted_index = verify_result.accepted_index if self.is_linear_att_mixed_model: linear_att_mtp_state_index_update( req_to_mtp_state_index=self.model.req_manager.req_to_mtp_state_index, @@ -530,7 +529,6 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): eagle_mem_indexes_cpu = self._draft_decode_func( model_input=model_input, - model_output=model_output, next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, mtp_accept_len=mtp_accept_len, @@ -550,17 +548,18 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() + self._update_mtp_verify_token_num(decode_reqs=decode_reqs) verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) # 第三阶段 event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() + self._update_mtp_accept_ratio(decode_reqs=decode_reqs, mtp_accept_len_cpu=mtp_accept_len_cpu) need_free_mem_indexes = model_input.mem_indexes_cpu[0:req_num][accepted_index_cpu == 0] if eagle_mem_indexes_cpu is not None: need_free_mem_indexes = torch.cat([need_free_mem_indexes, eagle_mem_indexes_cpu], dim=0) - self._update_mtp_accept_ratio(decode_reqs=decode_reqs, mtp_accept_len_cpu=mtp_accept_len_cpu) select_mask = torch.tensor(accepted_index_cpu, dtype=torch.bool, device="cpu") self._post_handle( run_reqs=verify_ok_reqs, @@ -584,7 +583,6 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): def _draft_decode_vanilla( self, model_input: ModelInput, - model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, mtp_accept_len: torch.Tensor, @@ -593,39 +591,40 @@ def _draft_decode_vanilla( all_next_token_ids = [] # share some inference info with the main model draft_model_input = model_input - draft_model_output = model_output - draft_next_token_ids_gpu = torch.zeros((model_input.batch_size), dtype=torch.int64, device="cuda") - if req_num > 0: - draft_next_token_ids_gpu[:req_num].copy_(next_token_ids, non_blocking=True) + draft_next_token_ids_gpu = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num > 0 else None, + batch_size=model_input.batch_size, + copy_len=req_num, + device=model_input.b_req_idx.device, + ) all_next_token_ids.append(draft_next_token_ids_gpu) # process the draft model output for draft_model_idx in range(self.mtp_step): - draft_model_input.input_ids = draft_next_token_ids_gpu - draft_model_input.mtp_draft_input_hiddens = draft_model_output.mtp_main_output_hiddens + draft_model_input = self.spec_adapter.prepare_draft_decode_input( + model_input=draft_model_input, + next_token_ids=draft_next_token_ids_gpu, + ) # spec decode: MTP draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) all_next_token_ids.append(draft_next_token_ids_gpu) if req_num > 0: - all_next_token_ids = torch.stack(all_next_token_ids, dim=1) # [batch_size, mtp_step + 1] - all_next_token_ids = all_next_token_ids[0:req_num, :] - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + self.spec_adapter.scatter_token_id_steps( + token_id_steps=all_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, b_req_idx=model_input.b_req_idx[:req_num], mtp_accept_len=mtp_accept_len, + row_count=req_num, ) return None def _draft_decode_eagle( self, model_input: ModelInput, - model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, mtp_accept_len: torch.Tensor, @@ -634,53 +633,50 @@ def _draft_decode_eagle( all_next_token_ids = [] # share some inference info with the main model draft_model_input = model_input - draft_model_output = model_output all_next_token_ids.append(next_token_ids) - draft_next_token_ids_gpu = torch.zeros((model_input.batch_size), dtype=torch.int64, device="cuda") - if req_num > 0: - draft_next_token_ids_gpu[:req_num].copy_(next_token_ids, non_blocking=True) + draft_next_token_ids_gpu = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num > 0 else None, + batch_size=model_input.batch_size, + copy_len=req_num, + device=model_input.b_req_idx.device, + ) real_req_num = req_num // (self.mtp_step + 1) padded_req_num = model_input.batch_size // (self.mtp_step + 1) - real_req_num - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(real_req_num * self.mtp_step) - eagle_mem_indexes_cpu = g_infer_context.req_manager.mem_manager.alloc(real_req_num * self.mtp_step) + eagle_mem_indexes_cpu = self.spec_adapter.alloc_extra_mem_indexes(real_req_num * self.mtp_step) eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) # process the draft model output for _step in range(self.mtp_step): - draft_model_input.input_ids = draft_next_token_ids_gpu - draft_model_input.mtp_draft_input_hiddens = draft_model_output.mtp_main_output_hiddens + draft_model_input = self.spec_adapter.prepare_draft_decode_input( + model_input=draft_model_input, + next_token_ids=draft_next_token_ids_gpu, + ) # spec decode: MTP draft_model_idx = _step % self.num_mtp_models draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) # update the meta info of the inference draft_model_input.b_seq_len += 1 draft_model_input.max_kv_seq_len += 1 - eagle_mem_indexes_i = eagle_mem_indexes[_step * real_req_num : (_step + 1) * real_req_num] - eagle_mem_indexes_i = F.pad( - input=eagle_mem_indexes_i, - pad=(0, padded_req_num), - mode="constant", - value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, + draft_model_input.mem_indexes = self.spec_adapter.append_padded_eagle_step_mem_indexes( + model_input=draft_model_input, + eagle_mem_indexes=eagle_mem_indexes, + step=_step, + real_req_num=real_req_num, + padded_req_num=padded_req_num, + mtp_step=self.mtp_step, ) - draft_model_input.mem_indexes = torch.cat( - [draft_model_input.mem_indexes.view(-1, self.mtp_step + 1)[:, 1:], eagle_mem_indexes_i.view(-1, 1)], - dim=1, - ).view(-1) draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) all_next_token_ids.append(draft_next_token_ids_gpu) if req_num > 0: - all_next_token_ids = torch.stack(all_next_token_ids, dim=1) # [batch_size, mtp_step + 1] - all_next_token_ids = all_next_token_ids[0:req_num, :] - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + self.spec_adapter.scatter_token_id_steps( + token_id_steps=all_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, b_req_idx=model_input.b_req_idx[:req_num], mtp_accept_len=mtp_accept_len, + row_count=req_num, ) return eagle_mem_indexes_cpu @@ -726,39 +722,27 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I b_prefill_has_output_cpu=b_has_out_cpu, ) - # spec prefill: MTP - draft_model_input0, draft_model_input1 = model_input0, model_input1 - draft_next_token_ids_gpu0 = torch.zeros((model_input0.batch_size), dtype=torch.int64, device="cuda") - if req_num0 > 0: - draft_next_token_ids_gpu0[0:req_num0].copy_(next_token_ids[0:req_num0], non_blocking=True) - - draft_next_token_ids_gpu1 = torch.zeros((model_input1.batch_size), dtype=torch.int64, device="cuda") - if req_num1 > 0: - draft_next_token_ids_gpu1[0:req_num1].copy_( - next_token_ids[req_num0 : (req_num0 + req_num1)], non_blocking=True - ) - - draft_model_output0, draft_model_output1 = model_output0, model_output1 - - for draft_model_idx in range(self.num_mtp_models): - - draft_model_input0 = prepare_mtp_prefill_inputs( - model_input=draft_model_input0, - b_next_token_ids=draft_next_token_ids_gpu0, - mtp_draft_input_hiddens=draft_model_output0.mtp_main_output_hiddens, - ) - - draft_model_input1 = prepare_mtp_prefill_inputs( - model_input=draft_model_input1, - b_next_token_ids=draft_next_token_ids_gpu1, - mtp_draft_input_hiddens=draft_model_output1.mtp_main_output_hiddens, - ) + draft_next_token_ids_gpu0 = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num0 > 0 else None, + batch_size=model_input0.batch_size, + copy_len=req_num0, + source_start=0, + device=model_input0.b_req_idx.device, + ) + draft_next_token_ids_gpu1 = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num1 > 0 else None, + batch_size=model_input1.batch_size, + copy_len=req_num1, + source_start=req_num0, + device=model_input1.b_req_idx.device, + ) - draft_model_output0, draft_model_output1 = self.draft_models[ - draft_model_idx - ].microbatch_overlap_prefill(draft_model_input0, draft_model_input1) - draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) - draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) + self.spec_adapter.build_initial_draft_state_overlap( + model_input0=model_input0, + next_token_ids0=draft_next_token_ids_gpu0, + model_input1=model_input1, + next_token_ids1=draft_next_token_ids_gpu1, + ) if req_num0 + req_num1 > 0 and g_infer_context.is_linear_att_mixed_model: _b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) @@ -833,23 +817,13 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf dtype=torch.int32, ).cuda(non_blocking=True) - mtp_accept_len, accepted_index = self._verify_mtp_v2( + verify_result = self.spec_adapter.verify_target_tokens( new_next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, ) - if self.is_linear_att_mixed_model: - b_mtp_index = torch.cat( - (model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0 - ) - linear_att_mtp_state_index_update( - req_to_mtp_state_index=self.model.req_manager.req_to_mtp_state_index, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=b_req_idx, - b_mtp_index=b_mtp_index, - accepted_index=accepted_index, - max_mtp_step=self.mtp_step + 1, - ) + mtp_accept_len = verify_result.accept_len + accepted_index = verify_result.accepted_index accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", gpu_tensor=accepted_index, @@ -866,8 +840,6 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf eagle_mem_indexes_cpu = self._draft_decode_overlap_func( model_input0=model_input0, model_input1=model_input1, - model_output0=model_output0, - model_output1=model_output1, b_req_idx=b_req_idx, next_token_ids=next_token_ids, mtp_accept_len=mtp_accept_len, @@ -888,11 +860,13 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if req_num0 + req_num1 > 0: event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() + self._update_mtp_verify_token_num(decode_reqs=decode_reqs) verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() + self._update_mtp_accept_ratio(decode_reqs=decode_reqs, mtp_accept_len_cpu=mtp_accept_len_cpu) mem_indexes_cpu = torch.cat( (model_input0.mem_indexes_cpu[0:req_num0], model_input1.mem_indexes_cpu[0:req_num1]), dim=0 ) @@ -900,7 +874,6 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if eagle_mem_indexes_cpu is not None: need_free_mem_indexes = torch.cat((need_free_mem_indexes, eagle_mem_indexes_cpu), dim=0) - self._update_mtp_accept_ratio(decode_reqs=decode_reqs, mtp_accept_len_cpu=mtp_accept_len_cpu) select_mask = torch.tensor(accepted_index_cpu, dtype=torch.bool, device="cpu") self._post_handle( run_reqs=verify_ok_reqs, @@ -919,27 +892,10 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf event_pack.notify_pre_post_handle() return - def _draft_prefill_forward(self, model_input: ModelInput, model_output: ModelOutput, next_token_ids: torch.Tensor): - # spec prefill: MTP, 这个地方只是为了填充draft model的 kv, 并不会使用生成的token_id。 - draft_model_input = model_input - draft_model_output = model_output - draft_next_token_ids_gpu = next_token_ids - for draft_model_idx in range(self.num_mtp_models): - draft_model_input = prepare_mtp_prefill_inputs( - model_input=draft_model_input, - b_next_token_ids=draft_next_token_ids_gpu, - mtp_draft_input_hiddens=draft_model_output.mtp_main_output_hiddens, - ) - draft_model_output = self.draft_models[draft_model_idx].forward(draft_model_input) - draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) - return - def _draft_decode_vanilla_overlap( self, model_input0: ModelInput, model_input1: ModelInput, - model_output0: ModelOutput, - model_output1: ModelOutput, b_req_idx: torch.Tensor, next_token_ids: torch.Tensor = None, mtp_accept_len: torch.Tensor = None, @@ -951,24 +907,35 @@ def _draft_decode_vanilla_overlap( all_next_token_ids.append(next_token_ids) # share some inference info with the main model draft_model_input0, draft_model_input1 = model_input0, model_input1 - draft_model_output0, draft_model_output1 = model_output0, model_output1 - - draft_next_token_ids_gpu0 = torch.zeros((model_input0.batch_size), dtype=torch.int64, device="cuda") - draft_next_token_ids_gpu1 = torch.zeros((model_input1.batch_size), dtype=torch.int64, device="cuda") - if req_num0 > 0: - draft_next_token_ids_gpu0[0:req_num0].copy_(next_token_ids[0:req_num0], non_blocking=True) - if req_num1 > 0: - draft_next_token_ids_gpu1[0:req_num1].copy_( - next_token_ids[req_num0 : (req_num0 + req_num1)], non_blocking=True - ) + + draft_next_token_ids_gpu0 = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num0 > 0 else None, + batch_size=model_input0.batch_size, + copy_len=req_num0, + source_start=0, + device=model_input0.b_req_idx.device, + ) + draft_next_token_ids_gpu1 = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num1 > 0 else None, + batch_size=model_input1.batch_size, + copy_len=req_num1, + source_start=req_num0, + device=model_input1.b_req_idx.device, + ) # process the draft model output for draft_model_idx in range(self.mtp_step): - draft_model_input0.input_ids = draft_next_token_ids_gpu0 - draft_model_input0.mtp_draft_input_hiddens = draft_model_output0.mtp_main_output_hiddens - draft_model_input1.input_ids = draft_next_token_ids_gpu1 - draft_model_input1.mtp_draft_input_hiddens = draft_model_output1.mtp_main_output_hiddens + draft_model_input0 = self.spec_adapter.prepare_draft_decode_input( + model_input=draft_model_input0, + next_token_ids=draft_next_token_ids_gpu0, + microbatch_index=0, + ) + draft_model_input1 = self.spec_adapter.prepare_draft_decode_input( + model_input=draft_model_input1, + next_token_ids=draft_next_token_ids_gpu1, + microbatch_index=1, + ) draft_model_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( draft_model_input0, draft_model_input1 @@ -982,11 +949,9 @@ def _draft_decode_vanilla_overlap( all_next_token_ids.append(draft_next_token_ids) if req_num0 + req_num1 > 0: - all_next_token_ids = torch.stack(all_next_token_ids, dim=1) - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + self.spec_adapter.scatter_token_id_steps( + token_id_steps=all_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, ) @@ -996,8 +961,6 @@ def _draft_decode_eagle_overlap( self, model_input0: ModelInput, model_input1: ModelInput, - model_output0: ModelOutput, - model_output1: ModelOutput, b_req_idx: torch.Tensor, next_token_ids: torch.Tensor = None, mtp_accept_len: torch.Tensor = None, @@ -1009,24 +972,27 @@ def _draft_decode_eagle_overlap( all_next_token_ids.append(next_token_ids) # share some inference info with the main model draft_model_input0, draft_model_input1 = model_input0, model_input1 - draft_model_output0, draft_model_output1 = model_output0, model_output1 - - draft_next_token_ids_gpu0 = torch.zeros((model_input0.batch_size), dtype=torch.int64, device="cuda") - draft_next_token_ids_gpu1 = torch.zeros((model_input1.batch_size), dtype=torch.int64, device="cuda") - if req_num0 > 0: - draft_next_token_ids_gpu0[0:req_num0].copy_(next_token_ids[0:req_num0], non_blocking=True) - if req_num1 > 0: - draft_next_token_ids_gpu1[0:req_num1].copy_( - next_token_ids[req_num0 : (req_num0 + req_num1)], non_blocking=True - ) + + draft_next_token_ids_gpu0 = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num0 > 0 else None, + batch_size=model_input0.batch_size, + copy_len=req_num0, + source_start=0, + device=model_input0.b_req_idx.device, + ) + draft_next_token_ids_gpu1 = self.spec_adapter.build_padded_next_token_ids( + token_ids=next_token_ids if req_num1 > 0 else None, + batch_size=model_input1.batch_size, + copy_len=req_num1, + source_start=req_num0, + device=model_input1.b_req_idx.device, + ) real_req_num0 = req_num0 // (self.mtp_step + 1) real_req_num1 = req_num1 // (self.mtp_step + 1) real_req_num = real_req_num0 + real_req_num1 padded_req_num0 = model_input0.batch_size // (self.mtp_step + 1) - real_req_num0 padded_req_num1 = model_input1.batch_size // (self.mtp_step + 1) - real_req_num1 - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(real_req_num * self.mtp_step) - eagle_mem_indexes_cpu = g_infer_context.req_manager.mem_manager.alloc(real_req_num * self.mtp_step) + eagle_mem_indexes_cpu = self.spec_adapter.alloc_extra_mem_indexes(real_req_num * self.mtp_step) eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) eagle_mem_indexes0 = eagle_mem_indexes[0 : real_req_num0 * self.mtp_step] eagle_mem_indexes1 = eagle_mem_indexes[real_req_num0 * self.mtp_step : real_req_num * self.mtp_step] @@ -1034,10 +1000,16 @@ def _draft_decode_eagle_overlap( # process the draft model output for _step in range(self.mtp_step): - draft_model_input0.input_ids = draft_next_token_ids_gpu0 - draft_model_input0.mtp_draft_input_hiddens = draft_model_output0.mtp_main_output_hiddens - draft_model_input1.input_ids = draft_next_token_ids_gpu1 - draft_model_input1.mtp_draft_input_hiddens = draft_model_output1.mtp_main_output_hiddens + draft_model_input0 = self.spec_adapter.prepare_draft_decode_input( + model_input=draft_model_input0, + next_token_ids=draft_next_token_ids_gpu0, + microbatch_index=0, + ) + draft_model_input1 = self.spec_adapter.prepare_draft_decode_input( + model_input=draft_model_input1, + next_token_ids=draft_next_token_ids_gpu1, + microbatch_index=1, + ) draft_model_idx = _step % self.num_mtp_models draft_model_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( @@ -1046,31 +1018,25 @@ def _draft_decode_eagle_overlap( draft_model_input0.b_seq_len += 1 draft_model_input0.max_kv_seq_len += 1 - eagle_mem_indexes_i = eagle_mem_indexes0[_step * real_req_num0 : (_step + 1) * real_req_num0] - eagle_mem_indexes_i = F.pad( - input=eagle_mem_indexes_i, - pad=(0, padded_req_num0), - mode="constant", - value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, + draft_model_input0.mem_indexes = self.spec_adapter.append_padded_eagle_step_mem_indexes( + model_input=draft_model_input0, + eagle_mem_indexes=eagle_mem_indexes0, + step=_step, + real_req_num=real_req_num0, + padded_req_num=padded_req_num0, + mtp_step=self.mtp_step, ) - draft_model_input0.mem_indexes = torch.cat( - [draft_model_input0.mem_indexes.view(-1, self.mtp_step + 1)[:, 1:], eagle_mem_indexes_i.view(-1, 1)], - dim=1, - ).view(-1) draft_model_input1.b_seq_len += 1 draft_model_input1.max_kv_seq_len += 1 - eagle_mem_indexes_i = eagle_mem_indexes1[_step * real_req_num1 : (_step + 1) * real_req_num1] - eagle_mem_indexes_i = F.pad( - input=eagle_mem_indexes_i, - pad=(0, padded_req_num1), - mode="constant", - value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, + draft_model_input1.mem_indexes = self.spec_adapter.append_padded_eagle_step_mem_indexes( + model_input=draft_model_input1, + eagle_mem_indexes=eagle_mem_indexes1, + step=_step, + real_req_num=real_req_num1, + padded_req_num=padded_req_num1, + mtp_step=self.mtp_step, ) - draft_model_input1.mem_indexes = torch.cat( - [draft_model_input1.mem_indexes.view(-1, self.mtp_step + 1)[:, 1:], eagle_mem_indexes_i.view(-1, 1)], - dim=1, - ).view(-1) draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) @@ -1080,11 +1046,9 @@ def _draft_decode_eagle_overlap( all_next_token_ids.append(draft_next_token_ids) if req_num0 + req_num1 > 0: - all_next_token_ids = torch.stack(all_next_token_ids, dim=1) - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + self.spec_adapter.scatter_token_id_steps( + token_id_steps=all_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py index 5b29ea0510..883f444e41 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py @@ -1,14 +1,28 @@ +from __future__ import annotations + import torch -from typing import List, Tuple +import triton +import triton.language as tl +from typing import TYPE_CHECKING, List, Tuple, Optional from lightllm.common.basemodel.triton_kernel.post_process.apply_penalty import apply_penalty from lightllm.common.basemodel.triton_kernel.post_process.apply_penalty_gpu_cache import apply_penalty_gpu_cache from lightllm.common.basemodel.triton_kernel.post_process.apply_invalid_token import apply_invalid_token_ids -from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.utils.envs_utils import get_env_start_args +if TYPE_CHECKING: + from lightllm.server.router.model_infer.infer_batch import InferReq + + +def sample( + logits: torch.Tensor, + reqs: List[InferReq], + eos_id: List[int] = [2], + dynamic_batch_size: Optional[int] = None, + selected_run_reqs: Optional[torch.Tensor] = None, +): + from lightllm.server.router.model_infer.infer_batch import g_infer_context + from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -def sample(logits: torch.Tensor, reqs: List[InferReq], eos_id: List[int] = [2]): ( b_req_idx, b_temperatures, @@ -24,6 +38,36 @@ def sample(logits: torch.Tensor, reqs: List[InferReq], eos_id: List[int] = [2]): skip_top_p, exist_req_use_random_seed, ) = _get_post_sample_tensors(reqs) + + sample_reqs = reqs + if selected_run_reqs is not None: + assert dynamic_batch_size is not None + ( + b_req_idx, + b_temperatures, + b_top_ps, + b_top_ks, + b_length_penalty_param, + b_mask_eos_reqs, + ) = _trim_post_sample_tensors( + dynamic_batch_size=dynamic_batch_size, + selected_run_reqs=selected_run_reqs, + b_req_idx=b_req_idx, + b_temperatures=b_temperatures, + b_top_ps=b_top_ps, + b_top_ks=b_top_ks, + b_length_penalty_param=b_length_penalty_param, + b_mask_eos_reqs=b_mask_eos_reqs, + ) + if has_invalid_token_ids or exist_req_use_random_seed: + sample_reqs = _get_selected_reqs(reqs=reqs, selected_run_reqs=selected_run_reqs) + if has_invalid_token_ids: + invalid_token_ids, cu_invalid_token_num, has_invalid_token_ids = _get_invalid_token_tensors( + reqs=sample_reqs + ) + if exist_req_use_random_seed: + exist_req_use_random_seed = any(req.generator is not None for req in sample_reqs) + eos_ids = g_pin_mem_manager.gen_from_list(key="eos_ids", data=eos_id, dtype=torch.int32).cuda(non_blocking=True) sampling_params_manager = g_infer_context.req_manager.req_sampling_params_manager @@ -85,13 +129,13 @@ def sample(logits: torch.Tensor, reqs: List[InferReq], eos_id: List[int] = [2]): elif skip_top_k and skip_top_p: # topk 等于整个词表,topp 等于1.0,等价于不进行topk topp过滤,直接进行随机采样,可以提升采样速度 - batch_next_token_ids = _random_sample(probs, reqs, exist_req_use_random_seed) + batch_next_token_ids = _random_sample(probs, sample_reqs, exist_req_use_random_seed) batch_next_token_probs = torch.gather(probs, dim=1, index=batch_next_token_ids.view(-1, 1)) return batch_next_token_ids.view(-1), torch.log(batch_next_token_probs).view(-1) else: batch_next_token_ids, batch_next_token_logprobs = _top_p_top_k_sample( - reqs, probs, b_top_ps, b_top_ks, exist_req_use_random_seed + sample_reqs, probs, b_top_ps, b_top_ks, exist_req_use_random_seed ) return batch_next_token_ids.view(-1), batch_next_token_logprobs.view(-1) @@ -155,6 +199,8 @@ def _random_sample(probs: torch.Tensor, reqs: List[InferReq], exist_req_use_rand def _get_post_sample_tensors(reqs: List[InferReq]): + from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + req_idxes: List[int] = [] temperatures: List[float] = [] top_ps: List[float] = [] @@ -231,3 +277,138 @@ def _get_post_sample_tensors(reqs: List[InferReq]): skip_top_p, exist_req_use_random_seed, ) + + +def _get_selected_reqs(reqs: List[InferReq], selected_run_reqs: torch.Tensor): + selected_run_reqs_cpu = selected_run_reqs.detach().cpu().tolist() + return [req_obj for req_obj, selected in zip(reqs, selected_run_reqs_cpu) if int(selected) != 0] + + +def _get_invalid_token_tensors(reqs: List[InferReq]): + from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + + invalid_token_ids: List[int] = [] + has_invalid_token_ids = False + cu_invalid_token_num = [0] + invalid_token_num_start = 0 + + for req_obj in reqs: + invalid_token_num_start += len(req_obj.sampling_param.invalid_token_ids) + cu_invalid_token_num.append(invalid_token_num_start) + if len(req_obj.sampling_param.invalid_token_ids) > 0: + has_invalid_token_ids = True + invalid_token_ids.extend(req_obj.sampling_param.invalid_token_ids) + + if not has_invalid_token_ids: + return None, None, False + + invalid_token_ids_cpu = g_pin_mem_manager.gen_from_list( + key="invalid_token_ids", data=invalid_token_ids, dtype=torch.int32 + ) + cu_invalid_token_num_cpu = g_pin_mem_manager.gen_from_list( + key="cu_invalid_token_num", data=cu_invalid_token_num, dtype=torch.int32 + ) + return ( + invalid_token_ids_cpu.cuda(non_blocking=True), + cu_invalid_token_num_cpu.cuda(non_blocking=True), + True, + ) + + +@triton.jit +def _fwd_kernel_trim_post_sample_tensors( + b_req_idx, + out_b_req_idx, + b_temperatures, + out_b_temperatures, + b_top_ps, + out_b_top_ps, + b_top_ks, + out_b_top_ks, + b_length_penalty_param, + out_b_length_penalty_param, + b_mask_eos_reqs, + out_b_mask_eos_reqs, + selected_run_reqs, + selected_dst_pos, + batch_size, + BLOCK_SIZE: tl.constexpr, +): + pid = tl.program_id(0) + offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size + selected = tl.load(selected_run_reqs + offsets, mask=mask, other=0) != 0 + dst_pos = tl.load(selected_dst_pos + offsets, mask=mask, other=0) + write_mask = mask & selected + + req_idx = tl.load(b_req_idx + offsets, mask=mask, other=0) + temperature = tl.load(b_temperatures + offsets, mask=mask, other=0.0) + top_p = tl.load(b_top_ps + offsets, mask=mask, other=0.0) + top_k = tl.load(b_top_ks + offsets, mask=mask, other=0) + length_penalty = tl.load(b_length_penalty_param + offsets, mask=mask, other=0) + mask_eos_req = tl.load(b_mask_eos_reqs + offsets, mask=mask, other=0) + + tl.store(out_b_req_idx + dst_pos, req_idx, mask=write_mask) + tl.store(out_b_temperatures + dst_pos, temperature, mask=write_mask) + tl.store(out_b_top_ps + dst_pos, top_p, mask=write_mask) + tl.store(out_b_top_ks + dst_pos, top_k, mask=write_mask) + tl.store(out_b_length_penalty_param + dst_pos, length_penalty, mask=write_mask) + tl.store(out_b_mask_eos_reqs + dst_pos, mask_eos_req, mask=write_mask) + + +def _trim_post_sample_tensors( + dynamic_batch_size: int, + selected_run_reqs: torch.Tensor, + b_req_idx: torch.Tensor, + b_temperatures: torch.Tensor, + b_top_ps: torch.Tensor, + b_top_ks: torch.Tensor, + b_length_penalty_param: torch.Tensor, + b_mask_eos_reqs: torch.Tensor, +): + assert selected_run_reqs.is_cuda + dynamic_batch_size = int(dynamic_batch_size) + selected_run_reqs = selected_run_reqs.to(torch.int32) + selected_dst_pos = torch.cumsum(selected_run_reqs, dim=0, dtype=torch.int32) - 1 + batch_size = selected_run_reqs.shape[0] + + out_b_req_idx = torch.empty((dynamic_batch_size,), dtype=b_req_idx.dtype, device=b_req_idx.device) + out_b_temperatures = torch.empty((dynamic_batch_size,), dtype=b_temperatures.dtype, device=b_temperatures.device) + out_b_top_ps = torch.empty((dynamic_batch_size,), dtype=b_top_ps.dtype, device=b_top_ps.device) + out_b_top_ks = torch.empty((dynamic_batch_size,), dtype=b_top_ks.dtype, device=b_top_ks.device) + out_b_length_penalty_param = torch.empty( + (dynamic_batch_size,), dtype=b_length_penalty_param.dtype, device=b_length_penalty_param.device + ) + out_b_mask_eos_reqs = torch.empty((dynamic_batch_size,), dtype=b_mask_eos_reqs.dtype, device=b_mask_eos_reqs.device) + + BLOCK_SIZE = 256 + grid = (triton.cdiv(batch_size, BLOCK_SIZE),) + _fwd_kernel_trim_post_sample_tensors[grid]( + b_req_idx=b_req_idx, + out_b_req_idx=out_b_req_idx, + b_temperatures=b_temperatures, + out_b_temperatures=out_b_temperatures, + b_top_ps=b_top_ps, + out_b_top_ps=out_b_top_ps, + b_top_ks=b_top_ks, + out_b_top_ks=out_b_top_ks, + b_length_penalty_param=b_length_penalty_param, + out_b_length_penalty_param=out_b_length_penalty_param, + b_mask_eos_reqs=b_mask_eos_reqs, + out_b_mask_eos_reqs=out_b_mask_eos_reqs, + selected_run_reqs=selected_run_reqs, + selected_dst_pos=selected_dst_pos, + batch_size=batch_size, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=4, + num_stages=1, + ) + + return ( + out_b_req_idx, + out_b_temperatures, + out_b_top_ps, + out_b_top_ks, + out_b_length_penalty_param, + out_b_mask_eos_reqs, + ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index ae294544ce..db7414e85a 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -5,7 +5,9 @@ from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.utils.envs_utils import ( enable_diverse_mode_gqa_decode_fast_kernel, + enable_triton_mtp_kernel, get_diverse_max_batch_shared_group_size, + enable_dynamic_mtp_verify, ) @@ -132,8 +134,13 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_mtp_index = torch.tensor(b_mtp_index, dtype=torch.int32, device="cpu") b_position_delta = build_b_position_delta(multimodal_params) + # diverse mode 和 dynamic MTP mode 使用不同的 shared group 构建逻辑 if enable_diverse_mode_gqa_decode_fast_kernel(): b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) + elif enable_dynamic_mtp_verify() or enable_triton_mtp_kernel(): + # MTP 模式下,使用专门的 shared group 构建函数 + b_shared_seq_len = None # MTP 模式不需要 b_shared_seq_len + b_mark_shared_group = build_mtp_shared_group_infos(run_reqs=run_reqs) else: b_shared_seq_len = None b_mark_shared_group = None @@ -217,3 +224,35 @@ def build_diverse_shared_group_infos(run_reqs: List[InferReq]) -> Tuple[torch.Te b_shared_seq_len = torch.tensor(b_shared_seq_len, dtype=torch.int32, device="cpu") b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cpu") return b_shared_seq_len, b_mark_shared_group + + +def build_mtp_shared_group_infos(run_reqs: List[InferReq]) -> torch.Tensor: + # Similar to build_diverse_shared_group_infos, + # but the grouping logic is based on b_mtp_index, which indicates the MTP step of each request + max_batch_shared_group_size = get_diverse_max_batch_shared_group_size() + req_ids = [req.req_id for req in run_reqs] + b_mark_shared_group = [] + _current_group = [] + for node in req_ids: + if not _current_group: + _current_group.append(node) + elif node == _current_group[-1]: + _current_group.append(node) + else: + b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) + b_mark_shared_group[-1] = len(_current_group) + _current_group.clear() + _current_group.append(node) + + if len(_current_group) == max_batch_shared_group_size: + b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) + b_mark_shared_group[-1] = len(_current_group) + _current_group.clear() + if _current_group: + b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) + b_mark_shared_group[-1] = len(_current_group) + _current_group.clear() + + assert len(b_mark_shared_group) == len(run_reqs) + b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cpu") + return b_mark_shared_group diff --git a/lightllm/server/router/model_infer/mode_backend/update_mem_index.py b/lightllm/server/router/model_infer/mode_backend/update_mem_index.py new file mode 100644 index 0000000000..4b0c7c441e --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/update_mem_index.py @@ -0,0 +1,50 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def update_eagle_mem_indexes_kernel( + old_mems_ptr, # [N] + new_step_mems_ptr, # [num_reqs] + b_req_mtp_start_loc, # [num_reqs] + out_mems_ptr, # [N] + req_all_num, + BLOCK_SIZE: tl.constexpr, +): + cur_req_idx = tl.program_id(0) + origin_req_num = tl.num_programs(0) + + offs = tl.arange(0, BLOCK_SIZE) + start_loc = tl.load(b_req_mtp_start_loc + cur_req_idx) + end_loc = tl.load(b_req_mtp_start_loc + cur_req_idx + 1, mask=cur_req_idx + 1 < origin_req_num, other=req_all_num) + + req_mtp_num = end_loc - start_loc + old_mems = tl.load(old_mems_ptr + start_loc + offs + 1, mask=offs + 1 < req_mtp_num, other=0) + tl.store(out_mems_ptr + start_loc + offs, old_mems, mask=offs + 1 < req_mtp_num) + new_step_mems = tl.load(new_step_mems_ptr + cur_req_idx) + tl.store(out_mems_ptr + end_loc - 1, new_step_mems) + + +def update_eagle_mem_indexes_triton( + old_mem_indexes: torch.Tensor, new_step_mem_indexes: torch.Tensor, b_req_mtp_start_loc: torch.Tensor +): + """ + old_mem_indexes: [N] CUDA Tensor + new_step_mem_indexes: [num_reqs] CUDA Tensor + """ + out = torch.empty_like(old_mem_indexes) + BLOCK_SIZE = 32 + original_num_reqs = b_req_mtp_start_loc.shape[0] + assert original_num_reqs == new_step_mem_indexes.shape[0] + req_all_num = old_mem_indexes.shape[0] + grid = (original_num_reqs,) + update_eagle_mem_indexes_kernel[grid]( + old_mems_ptr=old_mem_indexes, + new_step_mems_ptr=new_step_mem_indexes, + b_req_mtp_start_loc=b_req_mtp_start_loc, + out_mems_ptr=out, + req_all_num=req_all_num, + BLOCK_SIZE=BLOCK_SIZE, + ) + return out diff --git a/lightllm/server/router/model_infer/speculative/__init__.py b/lightllm/server/router/model_infer/speculative/__init__.py new file mode 100644 index 0000000000..6a238ecf00 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/__init__.py @@ -0,0 +1,42 @@ +__all__ = [ + "SpecRuntime", + "SpecDecodeForwardState", + "SpecDecodePostState", + "SpecDecodeRunner", + "SpecVerifier", + "SpecVerifyResult", + "build_spec_runtime", +] + + +def __getattr__(name): + if name in ("SpecRuntime", "build_spec_runtime"): + from lightllm.server.router.model_infer.speculative.runtime import SpecRuntime, build_spec_runtime + + values = { + "SpecRuntime": SpecRuntime, + "build_spec_runtime": build_spec_runtime, + } + return values[name] + if name in ("SpecDecodeForwardState", "SpecDecodePostState", "SpecDecodeRunner"): + from lightllm.server.router.model_infer.speculative.runner import ( + SpecDecodeForwardState, + SpecDecodePostState, + SpecDecodeRunner, + ) + + values = { + "SpecDecodeForwardState": SpecDecodeForwardState, + "SpecDecodePostState": SpecDecodePostState, + "SpecDecodeRunner": SpecDecodeRunner, + } + return values[name] + if name in ("SpecVerifier", "SpecVerifyResult"): + from lightllm.server.router.model_infer.speculative.verifier import SpecVerifier, SpecVerifyResult + + values = { + "SpecVerifier": SpecVerifier, + "SpecVerifyResult": SpecVerifyResult, + } + return values[name] + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/lightllm/server/router/model_infer/speculative/planner.py b/lightllm/server/router/model_infer/speculative/planner.py new file mode 100644 index 0000000000..21fd1bdbfa --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/planner.py @@ -0,0 +1,1626 @@ +from __future__ import annotations + +import math +import os +import random +from collections import Counter, deque +from dataclasses import dataclass +from typing import Dict, List, Optional, Tuple + +import numpy as np +from sortedcontainers import SortedDict + +from lightllm.utils.log_utils import init_logger + + +logger = init_logger(__name__) + + +@dataclass(frozen=True) +class SpecDecodePlan: + """Planner decision for one target decode iteration. + + Static MTP uses the full MTP-expanded target batch: + - dynamic_batch_size is None + - draft_step == mtp_step + + Dynamic MTP may compact target rows before forward: + - dynamic_batch_size is the selected target row count + - draft_step is the candidate length to generate after target verify + - pre_draft_step describes the previous iteration and controls whether + GPU verify sync can be skipped + """ + + dynamic_batch_size: Optional[int] + draft_step: int + pre_draft_step: int + selection_mode: str = "confidence" + + @property + def is_dynamic(self) -> bool: + return self.dynamic_batch_size is not None + + @property + def skip_verify_sync(self) -> bool: + return self.is_dynamic and self.pre_draft_step == 0 + + +class FixedMTPPlanner: + """Planner for static MTP.""" + + def __init__(self, mtp_step: int) -> None: + self.mtp_step = int(mtp_step) + + def plan(self, req_num: int | None = None, original_batch_size: int | None = None) -> SpecDecodePlan: + del req_num + del original_batch_size + return SpecDecodePlan( + dynamic_batch_size=None, + draft_step=self.mtp_step, + pre_draft_step=self.mtp_step, + selection_mode="none", + ) + + +class DynamicMTPPlanner: + planner_mode = "default" + + def __init__( + self, + mtp_step: int, + use_random_mode: bool = True, + random_mode_iter_threshold: int = 100, + ) -> None: + self.mtp_step = int(mtp_step) + + # 用于记录 decode 时的静态推理耗时(ms)。 + self.main_model_speeds = _InferCostMsTable() + self.draft_model_speeds = _InferCostMsTable() + + # 记录每个对应长度mtp step 步的接受概率。 由于原始位置必然是接受的,所以不需要记录。 + self.mtp_len_to_accept_ratio = [ + _EMAValue(decay=0.95, init_value=1.0, enable_decay_warmup=False) for _ in range(self.mtp_step) + ] + # 记录请求数量以及对应的推理dynamic_batch_size 对应的接受率统计 + self.req_num_to_dynamic_batch_size_to_accept_ratio: Dict[int, Dict[int, _EMAValue]] = {} + + # 每多少个请求采用随机的方式决定 dynamic_batch_size + self._iter = 0 + self._iter_threshold = int(random_mode_iter_threshold) + self._use_random_mode = bool(use_random_mode) + self._random = random.Random(0) + + # 记录上一次选择的draft step 步长,才好选择对应的 dynamic_batch_size + self.pre_draft_step = self.mtp_step + self._selection_mode = "confidence" + return + + def plan(self, req_num: int, original_batch_size: int) -> SpecDecodePlan: + dynamic_batch_size, draft_step, pre_draft_step = self.get_dynamic_batch_size( + req_num=req_num, + original_batch_size=original_batch_size, + ) + if dynamic_batch_size == original_batch_size and self._selection_mode == "confidence": + # Full-width dynamic plans do not need confidence sampling or + # tensor compaction, but ordinary calibration/controller plans + # still need to produce confidence for a potentially narrower + # next iteration. ``observe`` skips compaction while retaining + # those probabilities; the explicit profitable-chain path below + # uses ``full`` for the true Static-equivalent fast path. + self._selection_mode = "observe" + return SpecDecodePlan( + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + pre_draft_step=pre_draft_step, + selection_mode=self._selection_mode, + ) + + def update_infer_cost(self, *, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: + speed_table = self.draft_model_speeds if is_draft_model else self.main_model_speeds + speed_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + return + + def update_mtp_len_to_accept_ratio(self, mtp_len: int, accept_ratio: float) -> None: + assert mtp_len > 0 and mtp_len <= self.mtp_step + self.mtp_len_to_accept_ratio[mtp_len - 1].update(accept_ratio) + return + + def update_verified_prefix_stats(self, *, verify_len: int, accept_len: int) -> None: + if verify_len - 1 <= 0: + return + for mtp_index in range(verify_len - 1): + mtp_len = mtp_index + 1 + ratio = (accept_len - 1) / mtp_len + ratio = max(0.0, min(1.0, ratio)) + self.update_mtp_len_to_accept_ratio( + mtp_len=mtp_len, + accept_ratio=ratio, + ) + return + + def update_req_num_to_dynamic_batch_size_to_accept_ratio( + self, req_num: int, dynamic_batch_size: int, accept_ratio: float + ) -> None: + assert dynamic_batch_size >= req_num + self._get_req_num_to_dynamic_batch_size_to_accept_ratio( + req_num=req_num, dynamic_batch_size=dynamic_batch_size + ).update(accept_ratio) + return + + def _get_req_num_to_dynamic_batch_size_to_accept_ratio(self, req_num: int, dynamic_batch_size: int) -> "_EMAValue": + assert dynamic_batch_size >= req_num + if req_num not in self.req_num_to_dynamic_batch_size_to_accept_ratio: + self.req_num_to_dynamic_batch_size_to_accept_ratio[req_num] = {} + if dynamic_batch_size not in self.req_num_to_dynamic_batch_size_to_accept_ratio[req_num]: + self.req_num_to_dynamic_batch_size_to_accept_ratio[req_num][dynamic_batch_size] = _EMAValue( + decay=0.9, init_value=1.0, enable_decay_warmup=True + ) + return self.req_num_to_dynamic_batch_size_to_accept_ratio[req_num][dynamic_batch_size] + + def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: + """ + 返回 (dynamic_batch_size, draft_step, pre_draft_step)。 + pre_draft_step 是上一轮推理实际使用的 draft_step, + 调用方可据此判断当前 verify 是否有真实候选需要验证 + (pre_draft_step == 0 时 accept_len 恒为 1,无需等待 GPU verify 结果)。 + """ + assert req_num * (self.mtp_step + 1) == original_batch_size + pre_draft_step = self.pre_draft_step + if req_num == 0: + self.pre_draft_step = self.mtp_step + return 0, self.mtp_step, pre_draft_step + if not self.main_model_speeds.has_data() or not self.draft_model_speeds.has_data(): + # The cost model is only meaningful after both target and draft + # decode costs have been profiled. Block proposers such as DFlash + # do not run through draft_model.forward, and cudagraph may also be + # disabled, so a missing table must not collapse dynamic MTP to + # draft_step=0. + self.pre_draft_step = self.mtp_step + return req_num * (pre_draft_step + 1), self.mtp_step, pre_draft_step + + # case 1 如果采用随机的方式决定 dynamic_batch_size + self._iter += 1 + if self._use_random_mode and self._iter % self._iter_threshold == 0: + min_batch_size = req_num + max_batch_size = req_num * (pre_draft_step + 1) + dynamic_batch_size = self._random.randint(min_batch_size, max_batch_size) + + draft_step = self._random.randint(0, self.mtp_step) + self.pre_draft_step = draft_step + return dynamic_batch_size, draft_step, pre_draft_step + + # 通过计算的方式来获取已经知道的最优的 dynamic_batch_size,然后再决定 draft step 步长 + min_batch_size = req_num + max_batch_size = req_num * (pre_draft_step + 1) + dynamic_batch_size_keys = self.main_model_speeds.get_batch_size_keys_between(min_batch_size, max_batch_size) + + # 计算每个 dynamic_batch_size 对应的接受率以及单token的速度收益,然后选择最优的 dynamic_batch_size + cost_ms_list = [ + self._get_cost_ms(req_num=req_num, dynamic_batch_size=dynamic_batch_size, draft_step=pre_draft_step) + for dynamic_batch_size in dynamic_batch_size_keys + ] + dynamic_batch_size = dynamic_batch_size_keys[np.argmin(cost_ms_list)] + + # 下一步的 draft step 选择,需要考虑计算不同step步的收益问题再决定 + min_cost_ms = float("inf") + min_cost_ms_draft_step = 0 # 默认选择0步长 + for draft_step in range(0, self.mtp_step + 1): + cost_ms = self._get_cost_ms(req_num=req_num, dynamic_batch_size=dynamic_batch_size, draft_step=draft_step) + if cost_ms < min_cost_ms: + min_cost_ms = cost_ms + min_cost_ms_draft_step = draft_step + + # draft step 步长不能超过 mtp_step, 也不能小于0 + min_cost_ms_draft_step = min(min_cost_ms_draft_step, self.mtp_step) + min_cost_ms_draft_step = max(min_cost_ms_draft_step, 0) + self.pre_draft_step = min_cost_ms_draft_step + return dynamic_batch_size, min_cost_ms_draft_step, pre_draft_step + + def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: + accept_ratio = self._get_dynamic_batch_size_to_accept_ratio( + req_num=req_num, dynamic_batch_size=dynamic_batch_size + ) + total_time = ( + self.main_model_speeds.get(dynamic_batch_size) + + self.draft_model_speeds.get(dynamic_batch_size) * draft_step + ) + token_num = min((dynamic_batch_size * accept_ratio), req_num * (draft_step + 1)) + token_num = max(token_num, req_num) + cost_ms = total_time / token_num + return cost_ms + + def _get_dynamic_batch_size_to_accept_ratio(self, req_num: int, dynamic_batch_size: int): + ema = self._get_req_num_to_dynamic_batch_size_to_accept_ratio( + req_num=req_num, dynamic_batch_size=dynamic_batch_size + ) + if ema.get_count() >= 10: + # 当以及通过充分的数据统计以后,直接返回统计的接受率 + return ema.get() + + # 通过单请求的信息进行估计。 + real_step = dynamic_batch_size / req_num + assert real_step >= 1.0 + real_step = real_step - 1.0 + + # 用插值的方式估计不同mtp_len 对应的接受率 + left = int(math.floor(real_step)) + right = int(left + 1) + if left == 0: + left_value = 0.0 + else: + left_value = self.mtp_len_to_accept_ratio[left - 1].get() + + if right > self.mtp_step: + right_value = 0.0 + else: + right_value = self.mtp_len_to_accept_ratio[right - 1].get() + + accept_ratio = left_value + (right_value - left_value) * (real_step - left) + calcu_accept_ratio = (req_num + (dynamic_batch_size - req_num) * accept_ratio) / dynamic_batch_size + weight = ema.get_count() / 10 + # 通过统计数据和单请求数据进行加权平均,得到最终的接受率 + return calcu_accept_ratio * (1 - weight) + ema.get() * weight + + +class Eagle3DynamicMTPPlanner(DynamicMTPPlanner): + """Joint draft-length and verify-capacity planner for Eagle3. + + ``pre_draft_step`` bounds the proposal that is being verified now, while + ``draft_step`` controls the proposal built for the next target iteration. + Treating those as the same iteration makes ``draft_step == 0`` an + absorbing state. We instead choose the next draft length from a + steady-state search, then choose the current verify capacity within the + proposal width that is actually available now. + + Eagle3 also has an asymmetric draft cost: its first forward commits all K + selected target rows, while every recurrent forward after that processes + one accepted tail per request (batch B). + """ + + planner_mode = "eagle3" + _ACCEPT_RATIO_BUCKETS_PER_DRAFT_ROW = 8 + + def __init__(self, mtp_step: int) -> None: + super().__init__(mtp_step=mtp_step, use_random_mode=False) + # Eagle uses these values as a full-verify survival curve. Start from + # the first batch mean instead of decaying slowly from an all-accepted + # prior; otherwise a 32-iteration calibration still substantially + # overestimates short draft depths. + self._prefix_survival_decay = float(os.getenv("LIGHTLLM_EAGLE3_PREFIX_SURVIVAL_DECAY", "0.95")) + self.mtp_len_to_accept_ratio = [ + _EMAValue(decay=self._prefix_survival_decay, init_value=1.0, enable_decay_warmup=True) + for _ in range(self.mtp_step) + ] + self._min_static_progress_ratio = float(os.getenv("LIGHTLLM_EAGLE3_MIN_STATIC_PROGRESS_RATIO", "0.85")) + # This is the externally visible verify efficiency: + # accepted target rows / selected target verify rows. It includes the + # guaranteed first row, matching the project acceptance reported by + # the HTTP metrics and benchmark helper. + self._min_project_accept_ratio = float( + os.getenv( + "LIGHTLLM_EAGLE3_MIN_PROJECT_ACCEPT_RATIO", + # Keep a control margin above the externally requested 80% + # acceptance. Stop sequences can discard already-verified + # tail tokens, so HTTP output/verify metrics are about two + # points below the planner's model-accepted/verify feedback + # on GSM8K. + os.getenv("LIGHTLLM_EAGLE3_MIN_DRAFT_ACCEPT_RATIO", "0.860"), + ) + ) + self._full_verify_warmup_steps = max( + 0, + int(os.getenv("LIGHTLLM_EAGLE3_FULL_VERIFY_WARMUP_STEPS", "32")), + ) + self._full_verify_interval = max( + 0, + int(os.getenv("LIGHTLLM_EAGLE3_FULL_VERIFY_INTERVAL", "128")), + ) + self._early_full_probe_accept_ratio = float( + os.getenv("LIGHTLLM_EAGLE3_EARLY_FULL_PROBE_ACCEPT_RATIO", "0.72") + ) + self._early_full_probe_interval = max( + 1, + int(os.getenv("LIGHTLLM_EAGLE3_EARLY_FULL_PROBE_INTERVAL", "32")), + ) + self._progress_relax_ratio = float(os.getenv("LIGHTLLM_EAGLE3_PROGRESS_RELAX_RATIO", "1.0")) + self._capacity_accept_ratio_floor = float(os.getenv("LIGHTLLM_EAGLE3_CAPACITY_ACCEPT_RATIO_FLOOR", "0.80")) + self._capacity_feedback_gain = float(os.getenv("LIGHTLLM_EAGLE3_CAPACITY_FEEDBACK_GAIN", "0.10")) + self._capacity_feedback_reference_req_num = max( + 1, + int(os.getenv("LIGHTLLM_EAGLE3_CAPACITY_FEEDBACK_REFERENCE_REQ_NUM", "128")), + ) + self._align_verify_rows_to_graph = os.getenv("LIGHTLLM_EAGLE3_ALIGN_VERIFY_ROWS_TO_GRAPH", "1").lower() in { + "1", + "true", + "yes", + "on", + } + self._max_dynamic_draft_step = max( + 0, + min( + self.mtp_step, + int(os.getenv("LIGHTLLM_EAGLE3_MAX_DYNAMIC_DRAFT_STEP", str(self.mtp_step))), + ), + ) + self._three_regime_enabled = os.getenv("LIGHTLLM_EAGLE3_THREE_REGIME", "0").lower() in { + "1", + "true", + "yes", + "on", + } + self._adaptive_prefix_enabled = os.getenv("LIGHTLLM_EAGLE3_ADAPTIVE_PREFIX", "0").lower() in { + "1", + "true", + "yes", + "on", + } + self._adaptive_prefix_min_depth = max( + 0, + min( + self.mtp_step, + int(os.getenv("LIGHTLLM_EAGLE3_ADAPTIVE_PREFIX_MIN_DEPTH", "1")), + ), + ) + self._adaptive_prefix_cost_tolerance = max( + 0.0, + float(os.getenv("LIGHTLLM_EAGLE3_ADAPTIVE_PREFIX_COST_TOLERANCE", "0.02")), + ) + self._adaptive_prefix_long_chain_survival = float( + os.getenv("LIGHTLLM_EAGLE3_ADAPTIVE_PREFIX_LONG_CHAIN_SURVIVAL", "0.55") + ) + self._low_load_max_req_num = max( + 1, + int(os.getenv("LIGHTLLM_EAGLE3_LOW_LOAD_MAX_REQ_NUM", "16")), + ) + self._mid_load_max_req_num = max( + self._low_load_max_req_num, + int(os.getenv("LIGHTLLM_EAGLE3_MID_LOAD_MAX_REQ_NUM", "127")), + ) + self._low_load_prefix_depth = max( + 0, + int(os.getenv("LIGHTLLM_EAGLE3_LOW_LOAD_PREFIX_DEPTH", "7")), + ) + self._mid_load_prefix_depth = max( + 0, + int(os.getenv("LIGHTLLM_EAGLE3_MID_LOAD_PREFIX_DEPTH", "3")), + ) + self._active_req_num = 0 + self._full_verify_baseline_decay = float( + os.getenv( + "LIGHTLLM_EAGLE3_BATCH_ACCEPT_EMA_DECAY", + os.getenv("LIGHTLLM_EAGLE3_FULL_VERIFY_BASELINE_DECAY", "0.80"), + ) + ) + assert 0.0 < self._min_static_progress_ratio <= 1.0 + assert 0.0 < self._min_project_accept_ratio <= 1.0 + assert 0.0 < self._progress_relax_ratio <= 1.0 + assert 0.0 < self._capacity_accept_ratio_floor <= 1.0 + assert 0.0 < self._capacity_feedback_gain <= 1.0 + assert 0.0 <= self._early_full_probe_accept_ratio <= 1.0 + assert 0.0 <= self._prefix_survival_decay < 1.0 + assert 0.0 <= self._full_verify_baseline_decay < 1.0 + assert 0.0 <= self._adaptive_prefix_long_chain_survival <= 1.0 + + # Exact (B, K) statistics are sparse because the live request batch B + # changes constantly. Pool observations by normalized selected draft + # rows per request so adjacent concurrency levels share evidence. + self._accept_ratio_by_depth_and_width_bucket: Dict[Tuple[int, int], _EMAValue] = {} + self._accept_ratio_by_depth_req_and_batch: Dict[Tuple[int, int, int], _EMAValue] = {} + + # A few full-width verifies provide an unbiased estimate of the static + # Eagle acceptance length. Dynamic top-K observations alone are + # intentionally biased toward high-confidence rows and cannot serve as + # a static baseline. + self._full_verify_tokens_per_req_ema = _EMAValue( + decay=self._full_verify_baseline_decay, + init_value=float(self.mtp_step + 1), + enable_decay_warmup=True, + ) + self._full_verify_tokens_per_req_value = float(self.mtp_step + 1) + self._full_verify_accepted_token_sum = 0.0 + self._full_verify_request_count = 0 + self._full_verify_update_count = 0 + self._full_probe_pending = False + self._last_full_probe_plan_count = -self._early_full_probe_interval + + # This feedback comes from real non-full dynamic iterations. It + # provides a conservative capacity floor when a sparse cost-table + # candidate has an over-optimistic expected-token estimate. + self._observed_dynamic_draft_accept_ratio = _EMAValue( + decay=0.9, + init_value=0.75, + enable_decay_warmup=True, + ) + self._observed_dynamic_project_accept_ratio = _EMAValue( + decay=0.9, + init_value=self._min_project_accept_ratio, + enable_decay_warmup=True, + ) + self._observed_dynamic_tokens_per_req = _EMAValue( + decay=0.8, + init_value=1.0, + enable_decay_warmup=True, + ) + # Acceptance is aggregated with request weights. An iteration with a + # single long-tail request must not have the same influence as an + # iteration with 256 live requests. + self._observed_dynamic_accepted_token_sum = 0.0 + self._observed_dynamic_verify_row_sum = 0.0 + self._observed_dynamic_request_count = 0 + + # Closed-loop verify capacity. The initial value comes from the + # unbiased full-width baseline. Real dynamic iterations then move it + # toward the largest width that still satisfies project acceptance. + self._target_verify_rows_per_req_value: Optional[float] = None + + self._plan_log_interval = max(0, int(os.getenv("LIGHTLLM_EAGLE3_PLAN_LOG_INTERVAL", "0"))) + self._plan_count = 0 + self._draft_step_counts = Counter() + self._verify_rows_per_req_sum = 0.0 + self._expected_tokens_per_req_sum = 0.0 + + def update_verified_prefix_stats(self, *, verify_len: int, accept_len: int) -> None: + """Record the survival probability of each Eagle draft position. + + The generic planner records ``(accept_len - 1) / depth``. That value + is neither a conditional probability nor a survival probability and + can increase with depth. Eagle's expected-token model needs the + probability that a token at each depth is actually reached. + """ + + max_mtp_len = min(max(0, verify_len - 1), self.mtp_step) + for mtp_len in range(1, max_mtp_len + 1): + self.update_mtp_len_to_accept_ratio( + mtp_len=mtp_len, + accept_ratio=1.0 if accept_len > mtp_len else 0.0, + ) + + def update_verified_batch_prefix_stats( + self, + *, + verify_and_accept_lengths: List[Tuple[int, int]], + ) -> None: + """Update each depth once with a request-weighted batch mean.""" + + for mtp_len in range(1, self.mtp_step + 1): + eligible_accept_lengths = [ + accept_len for verify_len, accept_len in verify_and_accept_lengths if verify_len > mtp_len + ] + if not eligible_accept_lengths: + continue + survival_ratio = sum(accept_len > mtp_len for accept_len in eligible_accept_lengths) / len( + eligible_accept_lengths + ) + self.update_mtp_len_to_accept_ratio( + mtp_len=mtp_len, + accept_ratio=survival_ratio, + ) + + def update_req_num_to_dynamic_batch_size_to_accept_ratio( + self, + req_num: int, + dynamic_batch_size: int, + accept_ratio: float, + verify_step: int = None, + ) -> None: + depth = self.mtp_step if verify_step is None else int(verify_step) + exact_key = (depth, int(req_num), int(dynamic_batch_size)) + if exact_key not in self._accept_ratio_by_depth_req_and_batch: + self._accept_ratio_by_depth_req_and_batch[exact_key] = _EMAValue( + decay=0.9, + init_value=1.0, + enable_decay_warmup=True, + ) + self._accept_ratio_by_depth_req_and_batch[exact_key].update(accept_ratio) + self._get_width_bucket_accept_ratio( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=verify_step, + ).update(accept_ratio) + + def update_full_verify_tokens_per_req(self, tokens_per_req: float, req_num: int = 1) -> None: + tokens_per_req = max(1.0, min(float(self.mtp_step + 1), float(tokens_per_req))) + req_num = max(1, int(req_num)) + self._full_verify_accepted_token_sum += tokens_per_req * req_num + self._full_verify_request_count += req_num + # A lifetime cumulative average cannot follow a non-stationary trace: + # a preceding predictable segment would keep the progress floor high + # long after acceptance collapses. Track the live full-width baseline + # with an iteration-level EMA. Request counts remain cumulative only + # for diagnostics. + self._full_verify_tokens_per_req_ema.update(tokens_per_req) + self._full_verify_tokens_per_req_value = self._full_verify_tokens_per_req_ema.get() + self._full_verify_update_count += 1 + + def update_observed_iteration_stats( + self, + *, + tokens_per_req: float, + verify_rows_per_req: float, + is_full_verify: bool, + req_num: int = 1, + ) -> None: + if is_full_verify or verify_rows_per_req <= 1.0: + return + req_num = max(1, int(req_num)) + draft_accept_ratio = (tokens_per_req - 1.0) / (verify_rows_per_req - 1.0) + draft_accept_ratio = max(0.0, min(1.0, draft_accept_ratio)) + project_accept_ratio = max(0.0, min(1.0, tokens_per_req / verify_rows_per_req)) + self._observed_dynamic_draft_accept_ratio.update(draft_accept_ratio) + self._observed_dynamic_project_accept_ratio.update(project_accept_ratio) + self._observed_dynamic_tokens_per_req.update(tokens_per_req) + self._observed_dynamic_accepted_token_sum += tokens_per_req * req_num + self._observed_dynamic_verify_row_sum += verify_rows_per_req * req_num + self._observed_dynamic_request_count += req_num + + self._update_target_verify_rows_per_req( + tokens_per_req=tokens_per_req, + verify_rows_per_req=verify_rows_per_req, + req_num=req_num, + ) + + def _update_target_verify_rows_per_req( + self, + *, + tokens_per_req: float, + verify_rows_per_req: float, + req_num: int, + ) -> None: + current_target = self._get_controlled_verify_rows_per_req() + + # Holding the accepted-token count locally constant, this is the + # verify width that lands exactly on the project-acceptance target. + sample_target = tokens_per_req / self._min_project_accept_ratio + + # If progress is below its static-relative floor, acceptance and + # progress constraints conflict. Widen enough to recover progress; + # subsequent observations will pull the controller back once the + # additional rows stop paying for themselves. + min_expected_tokens = self._get_min_expected_tokens_per_req() + if self._get_observed_tokens_per_req() < min_expected_tokens: + sample_target = max( + sample_target, + verify_rows_per_req * min_expected_tokens / max(tokens_per_req, 1e-6), + ) + + sample_target = max(1.0, min(float(self.mtp_step + 1), sample_target)) + # Bound a single observation. This is especially important while a + # batch drains and only a handful of unusually hard requests remain. + sample_target = max(current_target - 0.5, min(current_target + 0.5, sample_target)) + request_weight = req_num / self._capacity_feedback_reference_req_num + alpha = 1.0 - (1.0 - self._capacity_feedback_gain) ** request_weight + self._target_verify_rows_per_req_value = current_target * (1.0 - alpha) + sample_target * alpha + + def _get_observed_project_accept_ratio(self) -> float: + if self._observed_dynamic_verify_row_sum <= 0.0: + return self._min_project_accept_ratio + return self._observed_dynamic_accepted_token_sum / self._observed_dynamic_verify_row_sum + + def _get_observed_draft_accept_ratio(self) -> float: + conditional_rows = self._observed_dynamic_verify_row_sum - self._observed_dynamic_request_count + if conditional_rows <= 0.0: + return self._observed_dynamic_draft_accept_ratio.get() + conditional_tokens = self._observed_dynamic_accepted_token_sum - self._observed_dynamic_request_count + return max(0.0, min(1.0, conditional_tokens / conditional_rows)) + + def _get_observed_tokens_per_req(self) -> float: + if self._observed_dynamic_request_count <= 0: + return self._observed_dynamic_tokens_per_req.get() + return self._observed_dynamic_accepted_token_sum / self._observed_dynamic_request_count + + def _get_dynamic_batch_size_to_accept_ratio( + self, + req_num: int, + dynamic_batch_size: int, + verify_step: int = None, + ): + depth = self.mtp_step if verify_step is None else int(verify_step) + exact_key = (depth, int(req_num), int(dynamic_batch_size)) + exact_ema = self._accept_ratio_by_depth_req_and_batch.get(exact_key) + base_estimate = super()._get_dynamic_batch_size_to_accept_ratio( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + ) + if exact_ema is not None and exact_ema.get_count() >= 10: + return exact_ema.get() + + width_ema = self._get_width_bucket_accept_ratio( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=verify_step, + ) + width_weight = min(1.0, width_ema.get_count() / 10.0) + return base_estimate * (1.0 - width_weight) + width_ema.get() * width_weight + + def _get_width_bucket_accept_ratio( + self, + *, + req_num: int, + dynamic_batch_size: int, + verify_step: int = None, + ) -> "_EMAValue": + selected_draft_rows_per_req = max(0.0, dynamic_batch_size / req_num - 1.0) + width_bucket = int(round(selected_draft_rows_per_req * self._ACCEPT_RATIO_BUCKETS_PER_DRAFT_ROW)) + depth = self.mtp_step if verify_step is None else int(verify_step) + bucket = (depth, width_bucket) + if bucket not in self._accept_ratio_by_depth_and_width_bucket: + self._accept_ratio_by_depth_and_width_bucket[bucket] = _EMAValue( + decay=0.9, + init_value=1.0, + enable_decay_warmup=True, + ) + return self._accept_ratio_by_depth_and_width_bucket[bucket] + + def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: + assert req_num * (self.mtp_step + 1) == original_batch_size + previous_req_num = self._active_req_num + self._active_req_num = int(req_num) + self._selection_mode = "confidence" + pre_draft_step = self.pre_draft_step + + if req_num == 0: + self.pre_draft_step = self.mtp_step + return 0, self.mtp_step, pre_draft_step + + max_batch_size = req_num * (pre_draft_step + 1) + if not self.main_model_speeds.has_data() or not self.draft_model_speeds.has_data(): + self._selection_mode = "observe" + self.pre_draft_step = self.mtp_step + return max_batch_size, self.mtp_step, pre_draft_step + + # Crossing from the prefix-friendly regime into high load invalidates + # the old per-request capacity target. Carrying that wide low-load + # target into a suddenly larger batch creates several expensive + # iterations before feedback contracts it. Clamp immediately to the + # high-load progress floor; normal closed-loop feedback can widen it + # again if the additional rows remain profitable. + crossed_into_high_load = ( + self._three_regime_enabled + and previous_req_num > 0 + and previous_req_num <= self._mid_load_max_req_num + and req_num > self._mid_load_max_req_num + ) + if crossed_into_high_load: + high_load_target = self._get_target_verify_rows_per_req() + if self._target_verify_rows_per_req_value is None: + self._target_verify_rows_per_req_value = high_load_target + else: + self._target_verify_rows_per_req_value = min( + self._target_verify_rows_per_req_value, + high_load_target, + ) + + # The legacy three-regime policy uses a fixed prefix depth. The + # adaptive variant below delays this decision until after full-probe + # handling, then chooses the prefix from the live survival curve and + # target+draft cost model. + prefix_depth = None if self._adaptive_prefix_enabled else self._get_load_prefix_depth(req_num=req_num) + if prefix_depth is not None: + self._selection_mode = "prefix" + verify_step = min(pre_draft_step, prefix_depth) + draft_step = min(self.mtp_step, prefix_depth) + dynamic_batch_size = req_num * (verify_step + 1) + self.pre_draft_step = draft_step + expected_tokens_per_req = ( + self._estimate_expected_token_num( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=pre_draft_step, + ) + / req_num + ) + self._record_plan_stats( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + expected_tokens_per_req=expected_tokens_per_req, + ) + return dynamic_batch_size, draft_step, pre_draft_step + + if self._should_schedule_full_probe(): + self._full_probe_pending = True + + force_full_verify = self._should_force_full_verify(pre_draft_step=pre_draft_step) + if force_full_verify: + self._selection_mode = "observe" + # Keep drafting full width during initial calibration. A periodic + # probe may return to the normal dynamic draft length immediately + # after its one full target verify. + draft_step = ( + self.mtp_step + if self._full_verify_update_count < self._full_verify_warmup_steps + else self._select_next_draft_step(req_num=req_num) + ) + dynamic_batch_size = max_batch_size + self.pre_draft_step = draft_step + expected_tokens_per_req = ( + self._estimate_expected_token_num( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=pre_draft_step, + ) + / req_num + ) + self._record_plan_stats( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + expected_tokens_per_req=expected_tokens_per_req, + ) + return dynamic_batch_size, draft_step, pre_draft_step + + # A highly predictable workload does not benefit from spending GPU + # work on confidence selection and row compaction. Keep every row of + # the current proposal and restore/retain a K-wide next proposal. On + # the following iteration ``DynamicMTPPlanner.plan`` marks the full + # width as a runtime fast path, so its target forward is identical to + # Static EAGLE3 while the planner continues monitoring acceptance. + # This rule is intentionally load-independent: high-concurrency, + # high-acceptance batches should expand just like low-load ones. + if ( + self._adaptive_prefix_enabled + and req_num <= self._mid_load_max_req_num + and self._has_profitable_long_chain() + ): + # If the previous proposal was shortened, only that prefix is + # valid even though ModelInput remains padded to the static width. + # Verify the valid prefix once while rebuilding K; the following + # iteration can then enter the true full-width fast path. + self._selection_mode = "full" if pre_draft_step == self.mtp_step else "prefix" + dynamic_batch_size = max_batch_size + draft_step = self.mtp_step + self.pre_draft_step = draft_step + expected_tokens_per_req = ( + self._estimate_expected_token_num( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=pre_draft_step, + ) + / req_num + ) + self._record_plan_stats( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + expected_tokens_per_req=expected_tokens_per_req, + ) + return dynamic_batch_size, draft_step, pre_draft_step + + adaptive_prefix_limit = self._get_adaptive_prefix_limit(req_num=req_num) + if adaptive_prefix_limit is not None: + self._selection_mode = "prefix" + prefix_depth = self._select_adaptive_prefix_depth( + req_num=req_num, + max_depth=adaptive_prefix_limit, + ) + verify_step = min(pre_draft_step, prefix_depth) + dynamic_batch_size = req_num * (verify_step + 1) + # A full probe needs two iterations after the proposal has been + # shortened: first rebuild a K-wide proposal, then verify it on + # the next target forward. Without this preparation step the + # pending probe can never observe deep positions and the planner + # becomes unable to recover a long chain after a workload shift. + draft_step = self.mtp_step if self._full_probe_pending else prefix_depth + self.pre_draft_step = draft_step + expected_tokens_per_req = ( + self._estimate_expected_token_num( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=pre_draft_step, + ) + / req_num + ) + self._record_plan_stats( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + expected_tokens_per_req=expected_tokens_per_req, + ) + return dynamic_batch_size, draft_step, pre_draft_step + + # Once the full-width baseline is calibrated, choose draft depth by + # the target+draft cost model, but control verify width with real + # project-acceptance feedback. Sparse (B, K) cost/acceptance buckets + # are too optimistic for unseen widths and previously widened GSM8K + # from about 4.5 to 6 rows/request before converging. + if self._full_verify_request_count > 0: + draft_step = self._select_next_draft_step(req_num=req_num) + if self._full_probe_pending: + draft_step = self.mtp_step + target_verify_rows_per_req = self._get_controlled_verify_rows_per_req() + dynamic_batch_size = int(math.ceil(req_num * target_verify_rows_per_req)) + if self._align_verify_rows_to_graph: + # Replay already pads an arbitrary target batch to the next + # captured shape. Turn those paid-for padding rows into real + # candidates so they can improve accepted progress for the + # same target-model graph cost. + graph_batch_size = self.main_model_speeds.get_ceil_batch_size( + dynamic_batch_size, + max_batch_size=max_batch_size, + ) + if graph_batch_size is not None: + dynamic_batch_size = graph_batch_size + dynamic_batch_size = min(max(dynamic_batch_size, req_num), max_batch_size) + self.pre_draft_step = draft_step + expected_tokens_per_req = ( + self._estimate_expected_token_num( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=pre_draft_step, + ) + / req_num + ) + self._record_plan_stats( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + expected_tokens_per_req=expected_tokens_per_req, + ) + return dynamic_batch_size, draft_step, pre_draft_step + + # The action selected here builds the proposal consumed by the next + # target iteration. Search it independently of the current proposal + # width so a zero-step iteration can recover on its own. + draft_step = self._select_next_draft_step(req_num=req_num) + if self._full_probe_pending: + draft_step = self.mtp_step + + dynamic_batch_size_keys = self._get_candidate_batch_sizes( + req_num=req_num, + max_batch_size=max_batch_size, + ) + candidates = [ + self._get_eagle3_candidate( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=pre_draft_step, + draft_step=draft_step, + ) + for dynamic_batch_size in dynamic_batch_size_keys + ] + best_candidate = self._select_best_candidate(candidates, req_num=req_num) + dynamic_batch_size = best_candidate[1] + expected_tokens_per_req = best_candidate[2] + self.pre_draft_step = draft_step + self._record_plan_stats( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + expected_tokens_per_req=expected_tokens_per_req, + ) + return dynamic_batch_size, draft_step, pre_draft_step + + def _get_target_verify_rows_per_req(self) -> float: + min_expected_tokens = self._get_min_expected_tokens_per_req() + if min_expected_tokens <= 1.0: + return 1.0 + return min( + float(self.mtp_step + 1), + min_expected_tokens / self._min_project_accept_ratio, + ) + + def _get_load_prefix_depth(self, *, req_num: int) -> Optional[int]: + if not self._three_regime_enabled: + return None + if req_num <= self._low_load_max_req_num: + return min(self.mtp_step, self._low_load_prefix_depth) + if req_num <= self._mid_load_max_req_num: + return min(self.mtp_step, self._mid_load_prefix_depth) + return None + + def _get_adaptive_prefix_limit(self, *, req_num: int) -> Optional[int]: + if not self._adaptive_prefix_enabled or not self._three_regime_enabled: + return None + if req_num <= self._low_load_max_req_num: + return min(self.mtp_step, self._low_load_prefix_depth) + if req_num <= self._mid_load_max_req_num: + return min(self.mtp_step, self._mid_load_prefix_depth) + return None + + def _select_adaptive_prefix_depth(self, *, req_num: int, max_depth: int) -> int: + """Choose a steady-state chain depth from measured prefix survival. + + Prefix selection is appropriate at low load because preserving a + request's chain can turn spare target capacity into forward progress. + The expected accepted length is computed directly from the survival + probability of every draft position, while the cost includes both the + target verification and Eagle3 proposal construction. If several + depths are effectively tied, prefer the longer chain so predictable + requests can retain near-K accepted lengths. + """ + + max_depth = max(self._adaptive_prefix_min_depth, min(self.mtp_step, int(max_depth))) + min_depth = min(self._adaptive_prefix_min_depth, max_depth) + if self._has_profitable_long_chain(max_depth=max_depth): + return max_depth + + candidates = [] + for depth in range(min_depth, max_depth + 1): + expected_tokens_per_req = 1.0 + sum( + self.mtp_len_to_accept_ratio[mtp_index].get() for mtp_index in range(depth) + ) + dynamic_batch_size = req_num * (depth + 1) + cost_ms = self._get_eagle3_cost_ms( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=depth, + draft_step=depth, + expected_token_num=req_num * expected_tokens_per_req, + ) + project_accept_ratio = expected_tokens_per_req / (depth + 1) + candidates.append((cost_ms, depth, project_accept_ratio)) + + feasible = [ + candidate for candidate in candidates if candidate[2] >= self._min_project_accept_ratio + ] + if not feasible: + best_accept_ratio = max(candidate[2] for candidate in candidates) + feasible = [ + candidate for candidate in candidates if candidate[2] >= best_accept_ratio - 1e-6 + ] + + best_cost = min(cost for cost, _, _ in feasible) + return max( + depth + for cost, depth, _ in feasible + if cost <= best_cost * (1.0 + self._adaptive_prefix_cost_tolerance) + ) + + def _has_profitable_long_chain(self, *, max_depth: int = None) -> bool: + max_depth = self.mtp_step if max_depth is None else min(self.mtp_step, int(max_depth)) + if max_depth <= 0: + return False + mean_survival = sum( + self.mtp_len_to_accept_ratio[mtp_index].get() for mtp_index in range(max_depth) + ) / max_depth + return mean_survival >= self._adaptive_prefix_long_chain_survival + + def _get_controlled_verify_rows_per_req(self) -> float: + if self._target_verify_rows_per_req_value is None: + return self._get_target_verify_rows_per_req() + return max( + 1.0, + min(float(self.mtp_step + 1), self._target_verify_rows_per_req_value), + ) + + def _get_candidate_batch_sizes(self, *, req_num: int, max_batch_size: int) -> List[int]: + candidates = set( + self.main_model_speeds.get_batch_size_keys_between( + req_num, + max_batch_size, + ) + ) + candidates.add(req_num) + candidates.add(max_batch_size) + exact_target = int(math.ceil(req_num * self._get_target_verify_rows_per_req())) + if req_num <= exact_target <= max_batch_size: + candidates.add(exact_target) + + # Add exact points along the feasible progress/acceptance frontier. + # CUDA graph timing keys are deliberately coarse; without these + # points the cost search can only choose between two widely separated + # capacities and often leaves useful acceptance headroom unused. + if self._full_verify_request_count > 0: + for progress_ratio in np.linspace(self._min_static_progress_ratio, 1.0, num=7): + expected_tokens_per_req = self._full_verify_tokens_per_req_value * float(progress_ratio) + verify_rows_per_req = expected_tokens_per_req / self._min_project_accept_ratio + frontier_batch_size = int(math.ceil(req_num * verify_rows_per_req)) + if req_num <= frontier_batch_size <= max_batch_size: + candidates.add(frontier_batch_size) + return sorted(candidates) + + def _should_schedule_full_probe(self) -> bool: + periodic_probe = ( + self._full_verify_interval > 0 + and self._full_verify_update_count >= self._full_verify_warmup_steps + and self._plan_count > 0 + and self._plan_count % self._full_verify_interval == 0 + ) + # At low/mid load, confidence-selected prefix rows expose a rising + # acceptance regime before the periodic unbiased K-wide probe arrives. + # Probe early once prefix acceptance is high enough, so the survival + # curve can discover newly profitable deep tokens and restore long + # chains without waiting tens of seconds. Keep a cooldown and disable + # this path at high load, where a full probe is materially expensive. + early_probe = ( + self._adaptive_prefix_enabled + and self._active_req_num > 0 + and self._active_req_num <= self._mid_load_max_req_num + and self._observed_dynamic_draft_accept_ratio.get() + >= self._early_full_probe_accept_ratio + and self._plan_count - self._last_full_probe_plan_count + >= self._early_full_probe_interval + and not self._has_profitable_long_chain() + ) + if periodic_probe or early_probe: + self._last_full_probe_plan_count = self._plan_count + return True + return False + + def _should_force_full_verify(self, *, pre_draft_step: int) -> bool: + if pre_draft_step != self.mtp_step: + return False + if self._full_verify_update_count < self._full_verify_warmup_steps: + return True + if self._full_probe_pending: + self._full_probe_pending = False + return True + return False + + def _record_plan_stats( + self, + *, + req_num: int, + dynamic_batch_size: int, + draft_step: int, + expected_tokens_per_req: float, + ) -> None: + self._plan_count += 1 + if self._plan_log_interval <= 0: + return + self._draft_step_counts[int(draft_step)] += 1 + self._verify_rows_per_req_sum += dynamic_batch_size / req_num + self._expected_tokens_per_req_sum += expected_tokens_per_req + if self._plan_count % self._plan_log_interval == 0: + logger.info( + "eagle3_dynamic_plan_stats plan_count=%d draft_step_counts=%s " + "avg_verify_rows_per_req=%.6f avg_expected_tokens_per_req=%.6f " + "static_tokens_per_req=%.6f full_verify_count=%d full_verify_req_count=%d " + "observed_project_accept_ratio=%.6f observed_draft_accept_ratio=%.6f " + "observed_tokens_per_req=%.6f target_verify_rows_per_req=%.6f " + "prefix_accept_ratios=%s", + self._plan_count, + dict(sorted(self._draft_step_counts.items())), + self._verify_rows_per_req_sum / self._plan_count, + self._expected_tokens_per_req_sum / self._plan_count, + self._full_verify_tokens_per_req_value, + self._full_verify_update_count, + self._full_verify_request_count, + self._get_observed_project_accept_ratio(), + self._get_observed_draft_accept_ratio(), + self._get_observed_tokens_per_req(), + self._get_controlled_verify_rows_per_req(), + [round(value.get(), 6) for value in self.mtp_len_to_accept_ratio], + ) + + def get_trace_stats(self) -> dict: + """Expose controller state for controlled scheduling ablations.""" + + return { + "batch_accept_len_estimate": self._full_verify_tokens_per_req_value, + "batch_accept_ema_decay": self._full_verify_baseline_decay, + "batch_accept_full_verify_updates": self._full_verify_update_count, + "target_verify_rows_per_req": self._get_controlled_verify_rows_per_req(), + "observed_tokens_per_req": self._get_observed_tokens_per_req(), + "min_expected_tokens_per_req": self._get_min_expected_tokens_per_req(), + } + + def _select_next_draft_step(self, *, req_num: int) -> int: + """Choose a recoverable long-run Eagle3 draft length. + + For each possible length, jointly search the verify capacities that + would be legal if that length were used repeatedly. This is a + one-state steady approximation to the next iteration and avoids + crediting newly drafted tokens to the current target forward. + """ + + candidates = [] + min_expected_tokens = self._get_min_expected_tokens_per_req() + for draft_step in range(self._max_dynamic_draft_step + 1): + # Even a full verify at this proposal depth must be capable of + # meeting the static-relative progress floor. This prevents an + # optimistic sparse (B, K) bucket from selecting draft_step=3 on + # GSM8K and paying for many extra target iterations. + max_depth_tokens_per_req = 1.0 + sum( + self.mtp_len_to_accept_ratio[mtp_index].get() for mtp_index in range(draft_step) + ) + if max_depth_tokens_per_req + 1e-6 < min_expected_tokens: + continue + max_batch_size = req_num * (draft_step + 1) + dynamic_batch_size_keys = self._get_candidate_batch_sizes( + req_num=req_num, + max_batch_size=max_batch_size, + ) + for dynamic_batch_size in dynamic_batch_size_keys: + candidates.append( + self._get_eagle3_candidate( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=draft_step, + draft_step=draft_step, + ) + ) + if not candidates: + return self._max_dynamic_draft_step + return self._select_best_candidate(candidates, req_num=req_num)[4] + + def _get_eagle3_candidate( + self, + *, + req_num: int, + dynamic_batch_size: int, + verify_step: int, + draft_step: int, + ) -> Tuple[float, int, float, float, int]: + expected_token_num = self._estimate_expected_token_num( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=verify_step, + ) + cost_ms = self._get_eagle3_cost_ms( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=verify_step, + draft_step=draft_step, + expected_token_num=expected_token_num, + ) + expected_tokens_per_req = expected_token_num / req_num + project_accept_ratio = max(0.0, min(1.0, expected_token_num / dynamic_batch_size)) + return ( + cost_ms, + dynamic_batch_size, + expected_tokens_per_req, + project_accept_ratio, + draft_step, + ) + + def _select_best_candidate( + self, + candidates: List[Tuple[float, int, float, float, int]], + *, + req_num: int, + ) -> Tuple[float, int, float, float, int]: + assert candidates + min_expected_tokens = self._get_min_expected_tokens_per_req() + min_verify_rows = self._get_min_verify_rows_per_req() + progress_candidates = [ + candidate + for candidate in candidates + if candidate[2] >= min_expected_tokens and candidate[1] / req_num >= min_verify_rows + ] + efficient_candidates = [ + candidate for candidate in progress_candidates if candidate[3] >= self._min_project_accept_ratio + ] + if efficient_candidates: + return min(efficient_candidates, key=lambda candidate: candidate[0]) + + # If noisy online estimates leave no candidate satisfying both hard + # constraints, minimize their worst relative violation. This avoids + # collapsing to the narrow acceptance-only plan or jumping to the + # wide progress-only plan while the exact boundary bucket converges. + drafted_candidates = [candidate for candidate in candidates if candidate[1] > req_num] + if drafted_candidates: + best_constraint_score = max( + min( + candidate[2] / max(min_expected_tokens, 1e-6), + candidate[3] / max(self._min_project_accept_ratio, 1e-6), + ) + for candidate in drafted_candidates + ) + balanced_candidates = [ + candidate + for candidate in drafted_candidates + if min( + candidate[2] / max(min_expected_tokens, 1e-6), + candidate[3] / max(self._min_project_accept_ratio, 1e-6), + ) + >= best_constraint_score - 1e-6 + ] + return min(balanced_candidates, key=lambda candidate: candidate[0]) + + # The current proposal may be too short to meet the floor after a + # transition or drafting may be disabled. Make maximal forward + # progress instead of falling back to a deceptively cheap iteration. + max_progress = max(candidate[2] for candidate in candidates) + max_progress_candidates = [candidate for candidate in candidates if candidate[2] >= max_progress - 1e-6] + return min(max_progress_candidates, key=lambda candidate: candidate[0]) + + def _get_min_expected_tokens_per_req(self) -> float: + if self._full_verify_request_count == 0: + return 1.0 + if self._three_regime_enabled and self._active_req_num > self._mid_load_max_req_num: + # Relax only the speculative gain above the guaranteed base row. + # Scaling the total accepted length can collapse the target below + # one token/request and makes the high-load controller degenerate + # into no-MTP. + static_gain = max(0.0, self._full_verify_tokens_per_req_value - 1.0) + return 1.0 + static_gain * self._min_static_progress_ratio * self._progress_relax_ratio + target_tokens = max( + 1.0, + self._full_verify_tokens_per_req_value * self._min_static_progress_ratio, + ) + if self._observed_dynamic_request_count > 0 and self._get_observed_tokens_per_req() >= target_tokens: + return max(1.0, target_tokens * self._progress_relax_ratio) + return target_tokens + + def _get_min_verify_rows_per_req(self) -> float: + min_expected_tokens = self._get_min_expected_tokens_per_req() + if min_expected_tokens <= 1.0: + return 1.0 + observed_project_accept_ratio = max( + self._capacity_accept_ratio_floor, + self._get_observed_project_accept_ratio(), + ) + return min_expected_tokens / observed_project_accept_ratio + + def _estimate_expected_token_num(self, *, req_num: int, dynamic_batch_size: int, verify_step: int) -> float: + accept_ratio = self._get_dynamic_batch_size_to_accept_ratio( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=verify_step, + ) + expected_token_num = min( + dynamic_batch_size * accept_ratio, + req_num * (verify_step + 1), + ) + return max(float(req_num), expected_token_num) + + def _get_eagle3_cost_ms( + self, + *, + req_num: int, + dynamic_batch_size: int, + verify_step: int, + draft_step: int, + expected_token_num: float = None, + ) -> float: + if expected_token_num is None: + expected_token_num = self._estimate_expected_token_num( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + verify_step=verify_step, + ) + + # Eagle3 commit runs on all selected verify rows. Every recurrent + # proposal step after that runs on one accepted tail per request. + draft_cost_ms = 0.0 + if draft_step > 0: + draft_cost_ms = self.draft_model_speeds.get(dynamic_batch_size) + if draft_step > 1: + draft_cost_ms += self.draft_model_speeds.get(req_num) * (draft_step - 1) + + total_time_ms = self.main_model_speeds.get(dynamic_batch_size) + draft_cost_ms + return total_time_ms / expected_token_num + + +class DSparkDynamicMTPPlanner(DynamicMTPPlanner): + """DSpark confidence-scheduled verify-capacity planner.""" + + planner_mode = "dspark" + + def __init__(self, mtp_step: int) -> None: + super().__init__( + mtp_step=mtp_step, + use_random_mode=False, + ) + self.mtp_len_to_accept_ratio = [ + _EMAValue(decay=0.95, init_value=1.0, enable_decay_warmup=True) for _ in range(self.mtp_step) + ] + self._predicted_dynamic_batch_sizes = deque(maxlen=2) + + def update_verified_prefix_stats(self, *, verify_len: int, accept_len: int) -> None: + if verify_len - 1 <= 0: + return + max_mtp_len = min(verify_len - 1, self.mtp_step) + for mtp_len in range(1, max_mtp_len + 1): + self.update_mtp_len_to_accept_ratio( + mtp_len=mtp_len, + accept_ratio=1.0 if accept_len > mtp_len else 0.0, + ) + return + + def update_predicted_schedule_probs(self, *, schedule_probs, req_num: int) -> None: + """Record a confidence-derived future capacity estimate. + + The current decode step still routes rows by the current probabilities + stored in req_to_next_token_probs. This queue is only used to choose + the future capacity K after a two-step delay, matching the asynchronous + scheduler constraint described by DSpark. + """ + + if req_num <= 0: + return + if not self.main_model_speeds.has_data(): + return + + probs = self._to_numpy(schedule_probs) + if probs is None or probs.ndim != 2 or probs.shape[1] <= 1: + return + + draft_probs = probs[:, 1 : self.mtp_step + 1] + if draft_probs.size == 0: + return + + valid_rows = np.any(draft_probs > 0.0, axis=1) + if not np.any(valid_rows): + return + + conditional_probs = np.clip(draft_probs[valid_rows], 0.01, 0.99) + survival_scores = np.cumprod(conditional_probs, axis=1) + dynamic_batch_size = self._select_dynamic_batch_size_from_survival_scores( + req_num=int(req_num), + survival_scores=survival_scores, + ) + self._predicted_dynamic_batch_sizes.append(dynamic_batch_size) + return + + def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: + assert req_num * (self.mtp_step + 1) == original_batch_size + pre_draft_step = self.pre_draft_step + self.pre_draft_step = self.mtp_step + if req_num == 0: + return 0, self.mtp_step, pre_draft_step + + max_batch_size = req_num * (pre_draft_step + 1) + if not self.main_model_speeds.has_data(): + return max_batch_size, self.mtp_step, pre_draft_step + + historical_batch_size = self._pop_historical_dynamic_batch_size( + req_num=req_num, + max_batch_size=max_batch_size, + ) + if historical_batch_size is not None: + return historical_batch_size, self.mtp_step, pre_draft_step + if len(self._predicted_dynamic_batch_sizes) > 0: + # A confidence estimate is available but has not satisfied the + # two-step async delay yet. Keep capacity conservative instead of + # leaking a same-step EMA fallback into DSpark scheduling. + return req_num, self.mtp_step, pre_draft_step + + candidate_batch_sizes = set(self.main_model_speeds.get_batch_size_keys_between(req_num, max_batch_size)) + candidate_batch_sizes.add(req_num) + candidate_batch_sizes.add(max_batch_size) + survival_prefix = self._estimate_survival_prefix(pre_draft_step) + + best_batch_size = req_num + best_throughput = -float("inf") + for candidate_batch_size in sorted(candidate_batch_sizes): + dynamic_batch_size = min(max(int(candidate_batch_size), req_num), max_batch_size) + expected_tokens = self._estimate_expected_tokens( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + survival_prefix=survival_prefix, + ) + verify_ms = max(self.main_model_speeds.get(dynamic_batch_size), 1e-6) + throughput = expected_tokens / verify_ms + if throughput > best_throughput: + best_throughput = throughput + best_batch_size = dynamic_batch_size + + return best_batch_size, self.mtp_step, pre_draft_step + + def _pop_historical_dynamic_batch_size(self, *, req_num: int, max_batch_size: int) -> Optional[int]: + if len(self._predicted_dynamic_batch_sizes) < 2: + return None + predicted_batch_size = int(self._predicted_dynamic_batch_sizes.popleft()) + return min(max(predicted_batch_size, int(req_num)), int(max_batch_size)) + + def _select_dynamic_batch_size_from_survival_scores( + self, + *, + req_num: int, + survival_scores: np.ndarray, + ) -> int: + flat_survival_scores = survival_scores.reshape(-1) + max_batch_size = int(req_num + flat_survival_scores.shape[0]) + if flat_survival_scores.shape[0] == 0: + return int(req_num) + + candidate_batch_sizes = set(self.main_model_speeds.get_batch_size_keys_between(req_num, max_batch_size)) + candidate_batch_sizes.add(int(req_num)) + candidate_batch_sizes.add(max_batch_size) + + candidate_batch_sizes = { + min(max(int(dynamic_batch_size), int(req_num)), max_batch_size) + for dynamic_batch_size in candidate_batch_sizes + } + selected_draft_counts = sorted( + {int(dynamic_batch_size) - int(req_num) for dynamic_batch_size in candidate_batch_sizes} + ) + expected_accepts_by_count = self._topk_prefix_sums( + values=flat_survival_scores, + counts=selected_draft_counts, + ) + + best_batch_size = int(req_num) + best_throughput = -float("inf") + for dynamic_batch_size in sorted(candidate_batch_sizes): + selected_draft_count = dynamic_batch_size - int(req_num) + expected_tokens = float(req_num) + float(expected_accepts_by_count[selected_draft_count]) + verify_ms = max(self.main_model_speeds.get(dynamic_batch_size), 1e-6) + throughput = expected_tokens / verify_ms + if throughput > best_throughput: + best_throughput = throughput + best_batch_size = dynamic_batch_size + return best_batch_size + + @staticmethod + def _topk_prefix_sums(*, values: np.ndarray, counts: List[int]) -> Dict[int, float]: + """Return sum(top-k(values)) only for the requested k values.""" + + if len(counts) == 0: + return {} + + flat_values = np.asarray(values, dtype=np.float64).reshape(-1) + value_count = int(flat_values.shape[0]) + ans: Dict[int, float] = {} + + normalized_counts = sorted({min(max(int(count), 0), value_count) for count in counts}) + if value_count == 0: + return {count: 0.0 for count in normalized_counts} + + partial_counts = [count for count in normalized_counts if 0 < count < value_count] + if len(partial_counts) > 0 and partial_counts[-1] > value_count // 2: + sorted_values = np.sort(flat_values)[::-1] + prefix_sums = np.cumsum(sorted_values, dtype=np.float64) + for count in normalized_counts: + ans[count] = 0.0 if count == 0 else float(prefix_sums[count - 1]) + return ans + + if normalized_counts[0] == 0: + ans[0] = 0.0 + if normalized_counts[-1] == value_count: + ans[value_count] = float(np.sum(flat_values, dtype=np.float64)) + if len(partial_counts) == 0: + return ans + + max_partial_count = partial_counts[-1] + top_values = np.partition(flat_values, value_count - max_partial_count)[value_count - max_partial_count :] + top_values.sort() + top_values = top_values[::-1] + prefix_sums = np.cumsum(top_values, dtype=np.float64) + for count in partial_counts: + ans[count] = float(prefix_sums[count - 1]) + return ans + + def _estimate_survival_prefix(self, pre_draft_step: int) -> List[float]: + survival_prefix = [1.0] + previous = 1.0 + for mtp_index in range(pre_draft_step): + survival = float(self.mtp_len_to_accept_ratio[mtp_index].get()) + survival = max(0.0, min(previous, survival)) + survival_prefix.append(survival) + previous = survival + return survival_prefix + + @staticmethod + def _estimate_expected_tokens( + *, + req_num: int, + dynamic_batch_size: int, + survival_prefix: List[float], + ) -> float: + expected_tokens = 0.0 + remaining = int(dynamic_batch_size) + for survival in survival_prefix: + take = min(req_num, remaining) + if take <= 0: + break + expected_tokens += take * float(survival) + remaining -= take + return max(float(req_num), expected_tokens) + + @staticmethod + def _to_numpy(value): + if value is None: + return None + if hasattr(value, "detach"): + value = value.detach().cpu().numpy() + return np.asarray(value, dtype=np.float64) + + +class _InferCostMsTable: + def __init__(self) -> None: + self.infer_cost_ms_table = SortedDict() + + def update(self, batch_size: int, infer_cost_ms: float) -> None: + assert batch_size > 0 + self.infer_cost_ms_table[int(batch_size)] = float(infer_cost_ms) + return + + def has_data(self) -> bool: + return len(self.infer_cost_ms_table) > 0 + + def get(self, batch_size: int) -> float: + assert batch_size > 0 + batch_size = int(batch_size) + + # 这种情况理论上不应该存在。 + if len(self.infer_cost_ms_table) == 0: + return batch_size * 1000.0 + + # 存在这个 batch_size 的记录,直接返回 + if batch_size in self.infer_cost_ms_table: + return self.infer_cost_ms_table[batch_size] + + # 不存在这个 batch_size,需要进行插值估计。 + max_batch_size = self.infer_cost_ms_table.peekitem(-1)[0] + if batch_size > max_batch_size: + max_infer_cost_ms = self.infer_cost_ms_table.peekitem(-1)[1] + # 这里面的 1000.0 意义是尽量使后续的估计,当超过最大graph支持的范围的时候,会直接倾向于关闭mtp功能。 + return max_infer_cost_ms + (batch_size - max_batch_size) * 1000.0 + else: + # 找到第一个大于等于 batch_size 的 key,并返回它的 value。 + index = self.infer_cost_ms_table.bisect_left(batch_size) + return self.infer_cost_ms_table.peekitem(index)[1] + + def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> List[int]: + assert batch_size1 > 0 and batch_size2 > 0 + start = min(int(batch_size1), int(batch_size2)) + end = max(int(batch_size1), int(batch_size2)) + ans = list(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) + if len(ans) == 0: + return [end] + else: + return ans + + def get_ceil_batch_size(self, batch_size: int, *, max_batch_size: int) -> Optional[int]: + """Return the next recorded graph shape without inventing a key. + + The cost table can be sparse or empty when CUDA graph is disabled, so + callers must be able to distinguish "no captured shape" from the + requested upper bound. + """ + + if len(self.infer_cost_ms_table) == 0: + return None + index = self.infer_cost_ms_table.bisect_left(int(batch_size)) + if index >= len(self.infer_cost_ms_table): + return None + candidate = int(self.infer_cost_ms_table.peekitem(index)[0]) + if candidate > int(max_batch_size): + return None + return candidate + + +class _EMAValue: + def __init__(self, decay: float, init_value: float, enable_decay_warmup: bool = True) -> None: + # decay=0 is the explicit latest-observation ablation used to compare + # EMA-smoothed online estimates against an otherwise identical + # controller. Production defaults remain strictly between zero/one. + assert 0.0 <= decay < 1.0 + self.enable_decay_warmup = enable_decay_warmup + self.decay = decay + + if self.enable_decay_warmup: + self.current_decay = 0.0 + else: + self.current_decay = self.decay + + self.value = init_value + self.second_moment_value = init_value ** 2 + self.update_count = 0 + + def update(self, new_value: float): + self.update_count += 1 + self.value = self.current_decay * self.value + (1.0 - self.current_decay) * new_value + self.second_moment_value = self.current_decay * self.second_moment_value + (1.0 - self.current_decay) * ( + new_value ** 2 + ) + # 更新 current_decay 的值,使得 current_decay 逐渐逼近 decay 的值 + self.current_decay = min(self.decay, (self.decay + self.current_decay) / 2.0 + 0.001) + return + + def get(self) -> float: + return self.value + + def get_count(self) -> int: + return self.update_count + + def get_second_moment(self) -> float: + return self.second_moment_value + + def get_variance(self) -> float: + return max(0.0, self.second_moment_value - self.value ** 2) + + def get_sigma(self) -> float: + return math.sqrt(self.get_variance()) + + +__all__ = [ + "DSparkDynamicMTPPlanner", + "DynamicMTPPlanner", + "Eagle3DynamicMTPPlanner", + "FixedMTPPlanner", + "SpecDecodePlan", +] diff --git a/lightllm/server/router/model_infer/speculative/proposers/__init__.py b/lightllm/server/router/model_infer/speculative/proposers/__init__.py new file mode 100644 index 0000000000..ab5e6fe7ba --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/proposers/__init__.py @@ -0,0 +1,61 @@ +from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal + + +def build_spec_proposer(runtime) -> BaseSpecProposer: + spec_config = runtime.spec_config + if spec_config.is_dspark: + from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer + + return DSparkProposer(runtime) + if spec_config.is_dflash: + from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer + + return DFlashProposer(runtime) + if spec_config.is_eagle3: + from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer + + return Eagle3Proposer(runtime) + if spec_config.uses_recurrent_draft_model: + from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer + + return EagleMTPProposer(runtime) + + from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer + + return VanillaMTPProposer(runtime) + + +def __getattr__(name): + if name == "DFlashProposer": + from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer + + return DFlashProposer + if name == "DSparkProposer": + from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer + + return DSparkProposer + if name == "EagleMTPProposer": + from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer + + return EagleMTPProposer + if name == "Eagle3Proposer": + from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer + + return Eagle3Proposer + if name == "VanillaMTPProposer": + from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer + + return VanillaMTPProposer + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + +__all__ = [ + "BaseSpecProposer", + "DFlashProposer", + "DSparkProposer", + "EagleMTPProposer", + "Eagle3Proposer", + "SpecProposal", + "VanillaMTPProposer", + "build_spec_proposer", +] diff --git a/lightllm/server/router/model_infer/speculative/proposers/base.py b/lightllm/server/router/model_infer/speculative/proposers/base.py new file mode 100644 index 0000000000..e27c543430 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/proposers/base.py @@ -0,0 +1,191 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, List, Optional, Union + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.speculative.runtime import SpecRuntime + + +@dataclass +class SpecProposal: + """Draft proposal returned by a proposer. + + `token_ids` is the LightLLM service equivalent of DeepSpec's + DraftProposal.verify_input_ids. It contains the target model's freshly + sampled token in column 0, followed by draft candidates: + + token_ids: [verify_batch, draft_step + 1] + + In static MTP, `draft_step == backend.mtp_step`. In dynamic MTP it may be + shorter and the runtime pads before scatter. + + `draft_probs` is intentionally narrower than DeepSpec's full + [B, K, vocab] probability tensor. Current LightLLM dynamic MTP only needs + the selected-token probability from each draft step: + + draft_probs[i]: [verify_batch] + + `schedule_probs` optionally overrides `draft_probs` for dynamic verify + row selection. It can be a list of per-step vectors or a dense + [verify_batch, draft_step] matrix. DSpark uses confidence-head conditional + acceptance probabilities here; runtime scatters them into the same + per-request buffer and the dynamic selector converts them to prefix + survival probabilities. + + `extra_mem_indexes_cpu` records draft-only KV slots allocated by recurrent + or block proposers. Eagle3 uses these slots for recurrent draft tokens; + DFlash uses them for current-block scratch query/mask K/V: + + extra_mem_indexes_cpu: [slot_count] + + """ + + token_ids: torch.Tensor + extra_mem_indexes_cpu: Optional[torch.Tensor] + draft_probs: Optional[List[torch.Tensor]] = None + schedule_probs: Optional[Union[List[torch.Tensor], torch.Tensor]] = None + # Actual rows processed by draft-model forwards. Recurrent Eagle can + # prune low-confidence deep chains, so this need not equal B * draft_step. + draft_forward_rows: Optional[int] = None + + +class BaseSpecProposer: + """Base class for algorithm-specific draft proposal generation. + + A proposer owns the draft-side state transition. The target model gives it + the current target token ids plus captured target hidden features through + SpecRuntime.prepare_draft_* methods. The proposer returns candidate ids + but does not verify acceptance; verification is handled by SpecVerifier. + """ + + def __init__(self, runtime: "SpecRuntime") -> None: + self.runtime = runtime + self.backend = runtime.backend + + @property + def enable_dynamic_mtp(self) -> bool: + return self.runtime.enable_dynamic_mtp + + def prepare_draft_prefill_input( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: Optional[torch.Tensor] = None, + microbatch_index: int = 0, + ) -> ModelInput: + return self.runtime.prepare_draft_prefill_input( + model_input=model_input, + next_token_ids=next_token_ids, + mtp_draft_input_hiddens=mtp_draft_input_hiddens, + microbatch_index=microbatch_index, + ) + + def prepare_draft_decode_input( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: Optional[torch.Tensor] = None, + microbatch_index: int = 0, + ) -> ModelInput: + return self.runtime.prepare_draft_decode_input( + model_input=model_input, + next_token_ids=next_token_ids, + mtp_draft_input_hiddens=mtp_draft_input_hiddens, + microbatch_index=microbatch_index, + ) + + def select_accepted_tail_rows(self, *, b_req_mtp_start_loc: torch.Tensor, accept_len: torch.Tensor) -> torch.Tensor: + return (b_req_mtp_start_loc + accept_len - 1).to(torch.long) + + def scatter_selected_step_probs( + self, + *, + selected_rows: torch.Tensor, + selected_probs: torch.Tensor, + verify_row_count: int, + ) -> torch.Tensor: + out = torch.zeros( + (verify_row_count,), + dtype=torch.float32, + device=selected_probs.device, + ) + out[selected_rows] = selected_probs.float() + return out + + def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: + """Allocate draft-owned temporary KV slots.""" + + return self.runtime.alloc_extra_mem_indexes(token_count) + + def build_initial_draft_state( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + ) -> None: + """Build initial draft KV/state before the first decode verify step. + + Inputs: + - `model_input`: target prompt ModelInput. Its request order and + mem_indexes are reused by the draft state builder. + - `next_token_ids`: first accepted target token, shape [run_req_num]. + + Runtime.prepare_draft_prefill_input injects captured target hidden + features into `mtp_draft_input_hiddens`. + + This hook only prepares draft-side state. It intentionally does not + scatter proposal tokens; the first decode iteration verifies as having + no draft candidates and produces the first proposal through + `propose_next`. + """ + + raise NotImplementedError + + def build_initial_draft_state_overlap( + self, + *, + model_input0: ModelInput, + next_token_ids0: torch.Tensor, + model_input1: ModelInput, + next_token_ids1: torch.Tensor, + ) -> None: + """Build initial draft state for two overlapped prefill microbatches.""" + + raise NotImplementedError + + def propose_next( + self, + *, + main_model_input: ModelInput, + main_model_output: Optional[ModelOutput] = None, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + verify_result=None, + ) -> SpecProposal: + """Generate candidate tokens after one target decode forward. + + Inputs: + - `main_model_input`: target decode ModelInput. In static MTP its + batch is laid out as [req0-main, req0-draft1, ...]. Dynamic MTP may + compact this batch before target forward. + - `next_token_ids`: target sampled ids for rows in `main_model_input`, + shape [verify_batch]. + - `b_req_mtp_start_loc`: start row for each logical request inside the + MTP-expanded batch, shape [logical_req_num]. + - `draft_step`: number of candidate draft tokens to produce. + - `verify_result`: optional target verification result from the just + finished target forward. Stateful block proposers use it to commit + the accepted target-hidden segment before preparing the next block. + + Returns a SpecProposal whose `token_ids[:, 0]` is `next_token_ids`. + """ + + raise NotImplementedError diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py new file mode 100644 index 0000000000..82b0c5c431 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -0,0 +1,189 @@ +from __future__ import annotations + +import copy + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal + + +class DFlashProposer(BaseSpecProposer): + """DFlash block proposer aligned to the Eagle3 runtime boundary. + + DFlash remains a non-causal block-prefill draft model, not a recurrent + token decoder. The service flow is: + - verify target tokens + - extend the DFlash draft KV cache with target hidden rows + - draft a new non-causal block from the accepted tail row + The memory lifecycle stays in the normal decode free path: + - rejected target token slots are freed by the normal decode free path + - current-block scratch KV uses extra mem slots returned through + `SpecProposal.extra_mem_indexes_cpu` + """ + + variant = "dflash" + + @torch.no_grad() + def build_initial_draft_state( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + ) -> None: + del next_token_ids + target_hidden = self.runtime.get_hidden() + if target_hidden.numel() == 0: + return + + draft_model = self.backend.draft_models[0] + assert model_input.input_ids is not None + assert model_input.input_ids.shape[0] == target_hidden.shape[0] + draft_input = copy.copy(model_input) + # DFlash consumes target hidden states directly on this prefill path. + draft_input.mtp_draft_input_hiddens = target_hidden + draft_model.forward(draft_input) + return + + @torch.no_grad() + def propose_next( + self, + *, + main_model_input: ModelInput, + main_model_output: ModelOutput = None, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + verify_result=None, + ) -> SpecProposal: + del main_model_output + assert 0 <= draft_step <= self.backend.mtp_step + assert verify_result is not None, "DFlash proposal requires target verify result" + + num_reqs = int(b_req_mtp_start_loc.shape[0]) + draft_model = self.backend.draft_models[0] + block_size = int(draft_model.block_size) + assert block_size >= draft_step + assert verify_result.accept_len.shape[0] == num_reqs + token_ids = next_token_ids.new_full( + (next_token_ids.shape[0], draft_step + 1), + fill_value=1, + ) + token_ids[:, 0] = next_token_ids + + if draft_step == 0: + return SpecProposal( + token_ids=token_ids, + extra_mem_indexes_cpu=None, + draft_probs=None, + ) + + self.extend_draft_kv_cache(main_model_input=main_model_input) + + # DFlash only drafts from the accepted tail row of each request. Unlike + # MTP, one anchor row expands to a whole non-causal block. + selected_rows = self.select_accepted_tail_rows( + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=verify_result.accept_len, + ) + draft_input, draft_mem_indexes_cpu = self.build_block_draft_input( + main_model_input=main_model_input, + next_token_ids=next_token_ids, + selected_rows=selected_rows, + num_reqs=num_reqs, + ) + draft_model_output = draft_model.forward(draft_input) + + flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + assert flat_token_ids.numel() == num_reqs * block_size + block_token_ids = flat_token_ids.reshape(num_reqs, block_size) + token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] + return SpecProposal( + token_ids=token_ids, + extra_mem_indexes_cpu=draft_mem_indexes_cpu, + draft_probs=None, + ) + + def extend_draft_kv_cache(self, *, main_model_input: ModelInput) -> None: + target_hidden = self.runtime.get_hidden() + draft_model = self.backend.draft_models[0] + batch_size = int(target_hidden.shape[0]) + + draft_kv_input = copy.copy(main_model_input) + draft_kv_input.batch_size = batch_size + draft_kv_input.total_token_num = batch_size + draft_kv_input.max_q_seq_len = 1 + draft_kv_input.prefix_total_token_num = 0 + draft_kv_input.is_prefill = True + # Each expanded MTP row writes one target-hidden KV slot for the same request. + draft_kv_input.b_ready_cache_len = main_model_input.b_seq_len - 1 + draft_kv_input.b_prefill_start_loc = torch.arange( + batch_size, + dtype=torch.int32, + device=target_hidden.device, + ) + draft_kv_input.b_position_delta = None + draft_kv_input.b_prefill_has_output_cpu = [False for _ in range(batch_size)] + draft_kv_input.mtp_draft_input_hiddens = target_hidden + draft_model.forward(draft_kv_input) + return + + def build_block_draft_input( + self, + *, + main_model_input: ModelInput, + next_token_ids: torch.Tensor, + selected_rows: torch.Tensor, + num_reqs: int, + ): + draft_model = self.backend.draft_models[0] + num_reqs = int(num_reqs) + block_size = int(draft_model.block_size) + assert selected_rows.shape[0] == num_reqs + draft_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * block_size) + + draft_input_ids = next_token_ids.new_full( + (num_reqs * block_size,), + fill_value=draft_model.mask_token_id, + ) + # Each block is [accepted_token, mask, ..., mask], matching DeepSpec's + # draft input. The proposer later maps the block logits back to + # [base_token + draft_tokens] for target verification. + draft_input_ids[::block_size] = next_token_ids.index_select(0, selected_rows) + + block_offsets = torch.arange( + block_size, + dtype=main_model_input.b_seq_len.dtype, + device=next_token_ids.device, + ) + draft_input = copy.copy(main_model_input) + draft_input.input_ids = draft_input_ids + draft_input.total_token_num = draft_input.input_ids.shape[0] + draft_input.batch_size = draft_input.total_token_num + draft_input.max_q_seq_len = 1 + draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size + draft_input.max_cache_len = draft_input.max_kv_seq_len + draft_input.b_req_idx = ( + main_model_input.b_req_idx.index_select(0, selected_rows).repeat_interleave(block_size).contiguous() + ) + draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) + # b_seq_len is real metadata, not cosmetic: copy_kv_index_to_req and + # FA3 use it to place scratch KV and compute the block cache length. + draft_input.b_seq_len = ( + (main_model_input.b_seq_len.index_select(0, selected_rows)[:, None] + block_offsets[None, :] + 1) + .reshape(-1) + .contiguous() + ) + if main_model_input.b_position_delta is not None: + draft_input.b_position_delta = ( + main_model_input.b_position_delta.index_select(0, selected_rows) + .repeat_interleave(block_size) + .contiguous() + ) + else: + draft_input.b_position_delta = torch.zeros_like(draft_input.b_req_idx) + draft_input.mem_indexes = draft_mem_indexes_cpu.cuda(non_blocking=True) + draft_input.b_mark_shared_group = torch.zeros_like(draft_input.b_req_idx) + draft_input.b_mark_shared_group[block_size - 1 :: block_size] = block_size + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(draft_input.batch_size)] + return draft_input, draft_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/speculative/proposers/dspark.py new file mode 100644 index 0000000000..4340e2754d --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/proposers/dspark.py @@ -0,0 +1,126 @@ +from __future__ import annotations + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer + + +class DSparkProposer(DFlashProposer): + """DSpark block proposer. + + DSpark shares DFlash's target-hidden KV injection and non-causal block + backbone. Its post layer returns Markov-corrected logits and optional + confidence logits, so the proposer follows the same token path as DFlash. + """ + + variant = "dspark" + + @torch.no_grad() + def propose_next( + self, + *, + main_model_input: ModelInput, + main_model_output: ModelOutput = None, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + verify_result=None, + ) -> SpecProposal: + del main_model_output + assert 0 <= draft_step <= self.backend.mtp_step + assert verify_result is not None, "DSpark proposal requires target verify result" + + num_reqs = int(b_req_mtp_start_loc.shape[0]) + verify_row_count = next_token_ids.shape[0] + draft_model = self.backend.draft_models[0] + block_size = int(draft_model.block_size) + assert block_size >= draft_step + assert verify_result.accept_len.shape[0] == num_reqs + + proposal_token_ids = next_token_ids.new_full( + (verify_row_count, draft_step + 1), + fill_value=1, + ) + proposal_token_ids[:, 0] = next_token_ids + schedule_probs = [] if self.enable_dynamic_mtp else None + + if draft_step == 0: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + draft_probs=[] if self.enable_dynamic_mtp else None, + schedule_probs=schedule_probs, + ) + + self.extend_draft_kv_cache(main_model_input=main_model_input) + selected_rows = self.select_accepted_tail_rows( + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=verify_result.accept_len, + ) + draft_input, draft_mem_indexes_cpu = self.build_block_draft_input( + main_model_input=main_model_input, + next_token_ids=next_token_ids, + selected_rows=selected_rows, + num_reqs=num_reqs, + ) + draft_model_output = draft_model.forward(draft_input) + + expected_block_rows = num_reqs * block_size + assert draft_model_output.logits.ndim >= 2, "draft logits must have a leading block-row dimension" + assert ( + draft_model_output.logits.shape[0] == expected_block_rows + ), f"draft logits rows must be {expected_block_rows}, got {draft_model_output.logits.shape[0]}" + flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + assert ( + flat_token_ids.numel() == expected_block_rows + ), f"draft token rows must be {expected_block_rows}, got {flat_token_ids.numel()}" + block_token_ids = flat_token_ids.reshape(num_reqs, block_size) + proposal_token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] + + draft_probs = None + if self.enable_dynamic_mtp: + confidence_logits = draft_model_output.mtp_draft_confidence_logits + if confidence_logits is None: + raise RuntimeError("DSpark dynamic verify requires confidence head logits") + assert confidence_logits.ndim == 2, "confidence logits must be [selected_rows, block_size]" + assert ( + confidence_logits.shape[0] == num_reqs + ), f"confidence logits rows must be {num_reqs}, got {confidence_logits.shape[0]}" + assert ( + confidence_logits.shape[1] >= draft_step + ), f"confidence logits columns must cover draft_step={draft_step}, got {confidence_logits.shape[1]}" + schedule_probs = self._scatter_step_probs( + selected_rows=selected_rows, + probs=confidence_logits[:, :draft_step].sigmoid(), + verify_row_count=verify_row_count, + ) + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=draft_mem_indexes_cpu, + draft_probs=draft_probs, + schedule_probs=schedule_probs, + ) + + def _scatter_step_probs( + self, + *, + selected_rows: torch.Tensor, + probs: torch.Tensor, + verify_row_count: int, + ): + assert selected_rows.ndim == 1, "selected_rows must be 1D" + assert probs.ndim == 2, "confidence probabilities must be [selected_rows, draft_step]" + assert probs.shape[0] == selected_rows.shape[0], ( + "confidence probability rows must match selected rows: " f"{selected_rows.shape[0]}, got {probs.shape[0]}" + ) + # Keep the async CPU capacity estimate aligned with the GPU dynamic + # selector, which clamps conditional draft probabilities before + # converting them to prefix survival scores. Unselected rows remain + # zero because scatter_selected_step_probs initializes the output. + probs = probs.clamp(min=0.01, max=0.99) + out = probs.new_zeros((verify_row_count, probs.shape[1]), dtype=torch.float32) + out[selected_rows, :] = probs.float() + return out diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py new file mode 100644 index 0000000000..f2be684106 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py @@ -0,0 +1,210 @@ +from __future__ import annotations + +import math +import os + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import RecurrentEagleMTPProposer + + +class Eagle3Proposer(RecurrentEagleMTPProposer): + """Eagle3 proposer. + + After target verification, Eagle3 commits the accepted target segment into + the draft cache with target hidden states, then drafts the next proposal + from that corrected state. + """ + + def __init__(self, runtime) -> None: + super().__init__(runtime) + self._confidence_draft_prune = os.getenv( + # Dynamic Eagle3 already ranks target verify rows by proposal + # confidence. Apply the same budget to deep draft frontiers by + # default; static Eagle3 is explicitly excluded below. + "LIGHTLLM_EAGLE3_CONFIDENCE_DRAFT_PRUNE", + "1", + ).lower() in {"1", "true", "yes", "on"} + self._draft_prune_safety_factor = max( + 0.0, + float(os.getenv("LIGHTLLM_EAGLE3_DRAFT_PRUNE_SAFETY_FACTOR", "1.10")), + ) + self._draft_prune_min_depth = max( + 2, + int(os.getenv("LIGHTLLM_EAGLE3_DRAFT_PRUNE_MIN_DEPTH", "4")), + ) + + def _get_pruned_active_count( + self, + *, + current_count: int, + draft_row_budget: int, + next_depth: int, + ) -> int: + if not self._confidence_draft_prune or next_depth < self._draft_prune_min_depth or current_count <= 1: + return current_count + + # A selected token at depth d consumes all d prefix draft rows from + # that request. Therefore at most L/d chains can reach depth d when + # the next target verify has L draft-row slots. Keep a configurable + # safety margin, then retain the highest-survival chain frontiers. + active_count = math.ceil(self._draft_prune_safety_factor * max(1, draft_row_budget) / next_depth) + return min(current_count, max(1, active_count)) + + def propose_next( + self, + *, + main_model_input: ModelInput, + main_model_output: ModelOutput = None, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + verify_result=None, + ) -> SpecProposal: + assert 0 <= draft_step <= self.backend.mtp_step + assert verify_result is not None, "Eagle3 proposal requires target verify result" + del main_model_output + verify_row_count = next_token_ids.shape[0] + num_reqs = b_req_mtp_start_loc.shape[0] + proposal_token_ids = next_token_ids.new_full( + (verify_row_count, draft_step + 1), + fill_value=1, + ) + proposal_token_ids[:, 0].copy_(next_token_ids) + collect_dynamic_probs = self.enable_dynamic_mtp and self.runtime.collect_dynamic_probs + draft_probs = [] if collect_dynamic_probs else None + + if draft_step == 0: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + draft_probs=draft_probs, + draft_forward_rows=0, + ) + + target_hidden = self.runtime.get_hidden() + + accept_len = verify_result.accept_len + # Scatter consumes the accepted-tail row for each request; only those + # rows need new draft columns after the commit step. + selected_rows = self.select_accepted_tail_rows( + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + ) + draft_model = self.backend.draft_models[0] + draft_model_input = self.make_full_verify_decode_input( + base_input=main_model_input, + input_ids=next_token_ids, + draft_hidden=target_hidden, + ) + draft_model_output = draft_model.forward(draft_model_input) + draft_logits = draft_model_output.logits.index_select(0, selected_rows) + if collect_dynamic_probs: + draft_next_token_ids, selected_draft_prob = self.backend._gen_argmax_token_ids_and_prob( + ModelOutput(logits=draft_logits) + ) + draft_prob = self.scatter_selected_step_probs( + selected_rows=selected_rows, + selected_probs=selected_draft_prob, + verify_row_count=verify_row_count, + ) + draft_probs.append(draft_prob) + chain_survival = selected_draft_prob.float().clamp(0.01, 0.99) + else: + draft_next_token_ids = self.backend._gen_argmax_token_ids(ModelOutput(logits=draft_logits)) + chain_survival = None + draft_hidden = self.runtime.get_hidden().index_select(0, selected_rows) + proposal_token_ids[selected_rows, 1] = draft_next_token_ids + draft_forward_rows = verify_row_count + + if draft_step == 1: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + draft_probs=draft_probs, + draft_forward_rows=draft_forward_rows, + ) + + eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * (draft_step - 1)) + eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) + + selected_seq_len = main_model_input.b_seq_len.index_select(0, selected_rows) + 1 + selected_req_idx = main_model_input.b_req_idx.index_select(0, selected_rows) + selected_mtp_index = torch.zeros_like(selected_req_idx) + selected_position_delta = ( + main_model_input.b_position_delta.index_select(0, selected_rows) + if main_model_input.b_position_delta is not None + else None + ) + one_row_group_marks = torch.ones((num_reqs,), dtype=torch.int32, device=next_token_ids.device) + draft_row_budget = max(1, int(main_model_input.batch_size) - int(num_reqs)) + + for step in range(1, draft_step): + next_depth = step + 1 + active_count = ( + self._get_pruned_active_count( + current_count=int(selected_rows.shape[0]), + draft_row_budget=draft_row_budget, + next_depth=next_depth, + ) + if collect_dynamic_probs + else int(selected_rows.shape[0]) + ) + if active_count < int(selected_rows.shape[0]): + assert chain_survival is not None + keep_rows = torch.topk( + chain_survival, + k=active_count, + largest=True, + sorted=False, + ).indices + selected_rows = selected_rows.index_select(0, keep_rows) + draft_next_token_ids = draft_next_token_ids.index_select(0, keep_rows) + draft_hidden = draft_hidden.index_select(0, keep_rows) + selected_seq_len = selected_seq_len.index_select(0, keep_rows) + selected_req_idx = selected_req_idx.index_select(0, keep_rows) + selected_mtp_index = selected_mtp_index.index_select(0, keep_rows) + if selected_position_delta is not None: + selected_position_delta = selected_position_delta.index_select(0, keep_rows) + one_row_group_marks = one_row_group_marks.index_select(0, keep_rows) + chain_survival = chain_survival.index_select(0, keep_rows) + + mem_start = (step - 1) * num_reqs + mem_indexes_i = eagle_mem_indexes[mem_start : mem_start + active_count] + draft_input = self.make_single_step_decode_input( + base_input=main_model_input, + input_ids=draft_next_token_ids, + draft_hidden=draft_hidden, + b_req_idx=selected_req_idx, + b_mtp_index=selected_mtp_index, + b_seq_len=selected_seq_len, + b_position_delta=selected_position_delta, + mem_indexes=mem_indexes_i, + b_mark_shared_group=one_row_group_marks, + max_kv_seq_len=main_model_input.max_kv_seq_len + step, + ) + draft_output = draft_model.forward(draft_input) + draft_forward_rows += active_count + if collect_dynamic_probs: + draft_next_token_ids, selected_draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_output) + draft_prob = self.scatter_selected_step_probs( + selected_rows=selected_rows, + selected_probs=selected_draft_prob, + verify_row_count=verify_row_count, + ) + draft_probs.append(draft_prob) + chain_survival = chain_survival * selected_draft_prob.float().clamp(0.01, 0.99) + else: + draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_output) + proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids + draft_hidden = self.runtime.get_hidden() + selected_seq_len = selected_seq_len + 1 + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=eagle_mem_indexes_cpu, + draft_probs=draft_probs, + draft_forward_rows=draft_forward_rows, + ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py new file mode 100644 index 0000000000..9efc6003e2 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -0,0 +1,219 @@ +from __future__ import annotations + +import copy + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer + + +class RecurrentEagleMTPProposer(VanillaMTPProposer): + """Shared draft-state setup for recurrent Eagle MTP proposers.""" + + def build_initial_draft_state( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + ) -> None: + draft_model_input = self.prepare_draft_prefill_input( + model_input=model_input, + next_token_ids=next_token_ids, + ) + self.backend.draft_models[0].forward(draft_model_input) + return None + + def build_initial_draft_state_overlap( + self, + *, + model_input0: ModelInput, + next_token_ids0: torch.Tensor, + model_input1: ModelInput, + next_token_ids1: torch.Tensor, + ) -> None: + draft_model_input0 = self.prepare_draft_prefill_input( + model_input=model_input0, + next_token_ids=next_token_ids0, + microbatch_index=0, + ) + draft_model_input1 = self.prepare_draft_prefill_input( + model_input=model_input1, + next_token_ids=next_token_ids1, + microbatch_index=1, + ) + self.backend.draft_models[0].microbatch_overlap_prefill(draft_model_input0, draft_model_input1) + return None + + def project_draft_decode_hidden(self, draft_hidden: torch.Tensor) -> torch.Tensor: + if draft_hidden is None: + return None + + draft_model = self.backend.draft_models[0] + pre_infer = getattr(draft_model, "pre_infer", None) + projector = getattr(pre_infer, "project_mtp_draft_hiddens", None) + if projector is None: + return draft_hidden + + # Keep draft decode CUDA graph input shape stable across target and draft hidden sources. + return projector( + draft_hidden, + draft_model.pre_post_weight, + use_custom_tensor_mananger=False, + ) + + def make_full_verify_decode_input( + self, + *, + base_input: ModelInput, + input_ids: torch.Tensor, + draft_hidden: torch.Tensor, + ) -> ModelInput: + new_input = copy.copy(base_input) + new_input.input_ids = input_ids + new_input.mtp_draft_input_hiddens = draft_hidden + new_input.mem_indexes_cpu = None + new_input.disable_mtp_decode_att = False + # The target may be using the fixed-layout full fast path. Draft + # forwards have their own graph/layout policy and must not inherit + # that target-only flag through the shallow ModelInput copy. + new_input.use_static_mtp_layout = False + return new_input + + def make_single_step_decode_input( + self, + *, + base_input: ModelInput, + input_ids: torch.Tensor, + draft_hidden: torch.Tensor, + b_req_idx: torch.Tensor, + b_mtp_index: torch.Tensor, + b_seq_len: torch.Tensor, + b_position_delta: torch.Tensor, + mem_indexes: torch.Tensor, + b_mark_shared_group: torch.Tensor, + max_kv_seq_len: int, + ) -> ModelInput: + new_input = copy.copy(base_input) + new_input.batch_size = b_seq_len.shape[0] + new_input.input_ids = input_ids + new_input.mtp_draft_input_hiddens = draft_hidden + new_input.b_req_idx = b_req_idx + new_input.b_mtp_index = b_mtp_index + new_input.b_seq_len = b_seq_len + new_input.b_position_delta = b_position_delta + new_input.mem_indexes = mem_indexes + new_input.mem_indexes_cpu = None + new_input.b_mark_shared_group = b_mark_shared_group + new_input.b_shared_seq_len = None + new_input.disable_mtp_decode_att = True + new_input.use_static_mtp_layout = False + new_input.max_q_seq_len = 1 + new_input.max_kv_seq_len = max_kv_seq_len + new_input.total_token_num = new_input.batch_size * max_kv_seq_len + # Recurrent Eagle decode only needs a correctly sized placeholder + # list. Nested per-row allocations otherwise sit between graph + # replays and extend the draft proposal critical path. + empty_multimodal_params = {"images": [], "audios": []} + new_input.multimodal_params = [empty_multimodal_params] * new_input.batch_size + return new_input + + +class EagleMTPProposer(RecurrentEagleMTPProposer): + """Recurrent Eagle MTP proposer. + + The draft model keeps a cache and repeatedly feeds the previous proposal + hidden back into the same draft model. + """ + + def propose_next( + self, + *, + main_model_input: ModelInput, + main_model_output: ModelOutput = None, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + verify_result=None, + ) -> SpecProposal: + assert 0 <= draft_step <= self.backend.mtp_step + del verify_result + return self._propose_expanded( + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=draft_step, + ) + + def _propose_expanded( + self, + *, + main_model_input: ModelInput, + main_model_output: ModelOutput = None, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + ) -> SpecProposal: + del main_model_output + num_reqs = b_req_mtp_start_loc.shape[0] + + if draft_step == 0: + eagle_mem_indexes_cpu = None + eagle_mem_indexes = None + else: + eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * draft_step) + eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) + + draft_model_input = main_model_input + draft_next_token_ids = next_token_ids + draft_hidden = self.runtime.get_hidden() if draft_step > 0 else None + all_next_token_ids = [next_token_ids] + draft_probs = [] if self.enable_dynamic_mtp else None + + for step in range(draft_step): + draft_hidden = self.project_draft_decode_hidden(draft_hidden) + draft_model_input = self.prepare_draft_decode_input( + model_input=draft_model_input, + next_token_ids=draft_next_token_ids, + mtp_draft_input_hiddens=draft_hidden, + ) + draft_model = self.backend.draft_models[0] + draft_model_output = draft_model.forward(draft_model_input) + draft_hidden = self.runtime.get_hidden() + + if self.enable_dynamic_mtp: + draft_next_token_ids, draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) + draft_probs.append(draft_prob) + else: + draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + + draft_model_input.b_seq_len += 1 + draft_model_input.max_kv_seq_len += 1 + eagle_mem_indexes_i = eagle_mem_indexes[step * num_reqs : (step + 1) * num_reqs] + if self.enable_dynamic_mtp: + from lightllm.server.router.model_infer.mode_backend.update_mem_index import ( + update_eagle_mem_indexes_triton, + ) + + draft_model_input.mem_indexes = update_eagle_mem_indexes_triton( + old_mem_indexes=draft_model_input.mem_indexes, + new_step_mem_indexes=eagle_mem_indexes_i, + b_req_mtp_start_loc=b_req_mtp_start_loc, + ) + else: + draft_model_input.mem_indexes = torch.cat( + [ + draft_model_input.mem_indexes.view(-1, self.backend.mtp_step + 1)[:, 1:], + eagle_mem_indexes_i.view(-1, 1), + ], + dim=1, + ).view(-1) + all_next_token_ids.append(draft_next_token_ids) + + return SpecProposal( + token_ids=torch.stack(all_next_token_ids, dim=1), + extra_mem_indexes_cpu=eagle_mem_indexes_cpu, + draft_probs=draft_probs, + ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py new file mode 100644 index 0000000000..4f5db86cdf --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py @@ -0,0 +1,119 @@ +from __future__ import annotations + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal + + +class VanillaMTPProposer(BaseSpecProposer): + """Chained MTP proposer. + + This path uses `mtp_step` independent draft modules. Step i consumes the + hidden feature produced by step i - 1 and predicts one candidate token. + + Target -> draft transfer: + - target prefill/decode captures final hidden states with shape + [token_num, hidden_size] + - runtime injects them into ModelInput.mtp_draft_input_hiddens before each + draft forward + """ + + def build_initial_draft_state( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + ) -> None: + draft_model_input = model_input + draft_next_token_ids = next_token_ids + draft_hidden = self.runtime.get_hidden() + for draft_model in self.backend.draft_models: + draft_model_input = self.prepare_draft_prefill_input( + model_input=draft_model_input, + next_token_ids=draft_next_token_ids, + mtp_draft_input_hiddens=draft_hidden, + ) + draft_model_output = draft_model.forward(draft_model_input) + draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + draft_hidden = self.runtime.get_hidden() + return None + + def build_initial_draft_state_overlap( + self, + *, + model_input0: ModelInput, + next_token_ids0: torch.Tensor, + model_input1: ModelInput, + next_token_ids1: torch.Tensor, + ) -> None: + draft_model_input0 = model_input0 + draft_model_input1 = model_input1 + draft_next_token_ids0 = next_token_ids0 + draft_next_token_ids1 = next_token_ids1 + draft_hidden0 = self.runtime.get_hidden(0) + draft_hidden1 = self.runtime.get_hidden(1) + + for draft_model in self.backend.draft_models: + draft_model_input0 = self.prepare_draft_prefill_input( + model_input=draft_model_input0, + next_token_ids=draft_next_token_ids0, + mtp_draft_input_hiddens=draft_hidden0, + microbatch_index=0, + ) + draft_model_input1 = self.prepare_draft_prefill_input( + model_input=draft_model_input1, + next_token_ids=draft_next_token_ids1, + mtp_draft_input_hiddens=draft_hidden1, + microbatch_index=1, + ) + draft_model_output0, draft_model_output1 = draft_model.microbatch_overlap_prefill( + draft_model_input0, + draft_model_input1, + ) + draft_next_token_ids0 = self.backend._gen_argmax_token_ids(draft_model_output0) + draft_next_token_ids1 = self.backend._gen_argmax_token_ids(draft_model_output1) + draft_hidden0 = self.runtime.get_hidden(0) + draft_hidden1 = self.runtime.get_hidden(1) + return None + + def propose_next( + self, + *, + main_model_input: ModelInput, + main_model_output: ModelOutput = None, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + verify_result=None, + ) -> SpecProposal: + del main_model_output + del verify_result + assert 0 <= draft_step <= self.backend.mtp_step + draft_model_input = main_model_input + draft_next_token_ids = next_token_ids + draft_hidden = self.runtime.get_hidden() if draft_step > 0 else None + all_next_token_ids = [next_token_ids] + draft_probs = [] if self.enable_dynamic_mtp else None + + for step in range(draft_step): + draft_model = self.backend.draft_models[step] + draft_model_input = self.prepare_draft_decode_input( + model_input=draft_model_input, + next_token_ids=draft_next_token_ids, + mtp_draft_input_hiddens=draft_hidden, + ) + draft_model_output = draft_model.forward(draft_model_input) + draft_hidden = self.runtime.get_hidden() + if self.enable_dynamic_mtp: + draft_next_token_ids, draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) + draft_probs.append(draft_prob) + else: + draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + all_next_token_ids.append(draft_next_token_ids) + + return SpecProposal( + token_ids=torch.stack(all_next_token_ids, dim=1), + extra_mem_indexes_cpu=None, + draft_probs=draft_probs, + ) diff --git a/lightllm/server/router/model_infer/speculative/runner.py b/lightllm/server/router/model_infer/speculative/runner.py new file mode 100644 index 0000000000..d3cc2ea882 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/runner.py @@ -0,0 +1,213 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Callable, List, Optional, Tuple + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.mtp_utils import ( + gen_b_req_mtp_start_loc, + linear_att_mtp_state_index_update, +) +from lightllm.server.router.model_infer.infer_batch import g_infer_context +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.speculative.planner import SpecDecodePlan + from lightllm.server.router.model_infer.speculative.runtime import SpecRuntime + + +@dataclass +class SpecDecodeForwardState: + model_input: ModelInput + original_run_reqs: List + plan: "SpecDecodePlan" + selected_run_reqs_cpu: Optional[torch.Tensor] + accepted_index_cpu: torch.Tensor + mtp_accept_len_cpu: torch.Tensor + next_token_ids_cpu: torch.Tensor + next_token_logprobs_cpu: torch.Tensor + next_token_ranks_cpu: torch.Tensor + verify_event: torch.cuda.Event + sync_event: torch.cuda.Event + additional_mem_indexes_cpu: Optional[torch.Tensor] + schedule_probs_cpu: Optional[torch.Tensor] + + +@dataclass +class SpecDecodePostState: + next_token_ids: torch.Tensor + next_token_logprobs: torch.Tensor + next_token_ranks: torch.Tensor + mtp_accept_len_cpu: torch.Tensor + need_free_mem_indexes: torch.Tensor + + +class SpecDecodeRunner: + def __init__(self, runtime: "SpecRuntime") -> None: + self.runtime = runtime + self.backend = runtime.backend + + def run_speculative_forward( + self, + *, + model_input: ModelInput, + model_output: ModelOutput, + run_reqs: List, + req_num: int, + plan: "SpecDecodePlan", + selected_run_reqs_cpu: Optional[torch.Tensor], + next_token_ids: torch.Tensor, + next_token_logprobs: torch.Tensor, + next_token_ranks: torch.Tensor, + copy_next_token_infos: Callable[ + [torch.Tensor, torch.Tensor, torch.Tensor], + Tuple[torch.Tensor, torch.Tensor, torch.Tensor], + ], + ) -> SpecDecodeForwardState: + runtime = self.runtime + b_req_mtp_start_loc = gen_b_req_mtp_start_loc(model_input.b_mtp_index, num_reqs=req_num) + verify_result = runtime.verify_target_tokens( + new_next_token_ids=next_token_ids, + b_req_idx=model_input.b_req_idx, + b_req_mtp_start_loc=b_req_mtp_start_loc, + ) + accepted_index = verify_result.accepted_index + if self.backend.is_linear_att_mixed_model: + linear_att_mtp_state_index_update( + req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, + accepted_index=accepted_index, + max_mtp_step=self.backend.mtp_step + 1, + ) + accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="accepted_index", + gpu_tensor=accepted_index, + ) + + verify_event = torch.cuda.Event(enable_timing=True) + verify_event.record() + + runtime.configure_dynamic_prob_collection(plan=plan) + proposal = runtime.propose_next( + main_model_input=model_input, + main_model_output=model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=plan.draft_step, + verify_result=verify_result, + ) + + all_next_token_ids = runtime.pad_all_next_token_ids( + token_ids=proposal.token_ids, + draft_step=plan.draft_step, + ) + all_next_token_probs = runtime.build_all_next_token_probs( + next_token_logprobs=next_token_logprobs, + proposal=proposal, + draft_step=plan.draft_step, + ) + schedule_probs_cpu = ( + g_pin_mem_manager.async_copy_from_gpu_tensor( + key="mtp_schedule_probs", + gpu_tensor=all_next_token_probs, + ) + if all_next_token_probs is not None and runtime.needs_schedule_probs_cpu() + else None + ) + + runtime.scatter_next_tokens( + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=model_input.b_req_idx, + mtp_accept_len=verify_result.accept_len, + all_next_token_probs=all_next_token_probs, + ) + + next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu = copy_next_token_infos( + next_token_ids, + next_token_logprobs, + next_token_ranks, + ) + + g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( + b_req_idx=model_input.b_req_idx, + next_token_ids=next_token_ids, + mask=accepted_index == 1, + ) + + mtp_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="mtp_accept_len", + gpu_tensor=verify_result.accept_len, + ) + + sync_event = torch.cuda.Event() + sync_event.record() + + return SpecDecodeForwardState( + model_input=model_input, + original_run_reqs=run_reqs, + plan=plan, + selected_run_reqs_cpu=selected_run_reqs_cpu, + accepted_index_cpu=accepted_index_cpu, + mtp_accept_len_cpu=mtp_accept_len_cpu, + next_token_ids_cpu=next_token_ids_cpu, + next_token_logprobs_cpu=next_token_logprobs_cpu, + next_token_ranks_cpu=next_token_ranks_cpu, + verify_event=verify_event, + sync_event=sync_event, + additional_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, + schedule_probs_cpu=schedule_probs_cpu, + ) + + def resolve_pre_post_reqs(self, *, state: SpecDecodeForwardState, decode_reqs: List): + if state.plan.skip_verify_sync: + assert self.runtime.enable_dynamic_mtp, "skip_verify_sync should only be True when dynamic MTP is enabled" + return decode_reqs, decode_reqs + + state.verify_event.synchronize() + return self.runtime.build_decode_req_lists( + original_run_reqs=state.original_run_reqs, + selected_run_reqs_cpu=state.selected_run_reqs_cpu, + accepted_index_cpu=state.accepted_index_cpu, + ) + + def finish_post(self, *, state: SpecDecodeForwardState, req_num: int, run_reqs: List) -> SpecDecodePostState: + state.sync_event.synchronize() + + runtime = self.runtime + if runtime.enable_dynamic_mtp: + runtime.update_dynamic_accept_stats( + req_num=req_num, + run_reqs=run_reqs, + accepted_index_cpu=state.accepted_index_cpu, + mtp_accept_len_cpu=state.mtp_accept_len_cpu, + dynamic_batch_size=state.plan.dynamic_batch_size, + verify_step=state.plan.pre_draft_step, + selection_mode=state.plan.selection_mode, + ) + if state.schedule_probs_cpu is not None: + runtime.update_dynamic_schedule_stats( + req_num=req_num, + schedule_probs_cpu=state.schedule_probs_cpu, + ) + + need_free_mem_indexes = runtime.build_decode_free_mem_indexes_cpu( + model_input=state.model_input, + selected_run_reqs_cpu=state.selected_run_reqs_cpu, + accepted_index_cpu=state.accepted_index_cpu, + ) + if state.additional_mem_indexes_cpu is not None: + need_free_mem_indexes = torch.cat([need_free_mem_indexes, state.additional_mem_indexes_cpu], dim=0) + + select_mask = state.accepted_index_cpu.to(dtype=torch.bool) + return SpecDecodePostState( + next_token_ids=state.next_token_ids_cpu[select_mask], + next_token_logprobs=state.next_token_logprobs_cpu[select_mask], + next_token_ranks=state.next_token_ranks_cpu[select_mask], + mtp_accept_len_cpu=state.mtp_accept_len_cpu, + need_free_mem_indexes=need_free_mem_indexes, + ) diff --git a/lightllm/server/router/model_infer/speculative/runtime.py b/lightllm/server/router/model_infer/speculative/runtime.py new file mode 100644 index 0000000000..4869109711 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/runtime.py @@ -0,0 +1,774 @@ +from __future__ import annotations + +import json +import os +from typing import Callable, List, Optional, Tuple + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.speculative.config import SpeculativeConfig +from lightllm.server.router.model_infer.speculative.planner import FixedMTPPlanner, SpecDecodePlan +from lightllm.server.router.model_infer.speculative.proposers import build_spec_proposer +from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.speculative.runner import ( + SpecDecodeForwardState, + SpecDecodePostState, + SpecDecodeRunner, +) +from lightllm.server.router.model_infer.speculative.state import SpecForwardContext, SpecHiddenStore +from lightllm.server.router.model_infer.speculative.verifier import SpecVerifier, SpecVerifyResult + + +class SpecRuntime: + """Facade between LightLLM backend code and speculative algorithms. + + The runtime keeps speculative decoding out of BaseModel and + chunked_prefill: + - BaseModel only calls hidden-capture methods. + - proposer implementations own draft-model state/proposal generation. + - SpecVerifier owns service-specific verify/scatter kernels. + + The main target->draft data path is: + 1. target model forward captures hidden features into SpecHiddenStore + 2. runtime injects those features into ModelInput.mtp_draft_input_hiddens + 3. proposer forwards the draft model and returns SpecProposal.token_ids + 4. verifier checks target acceptance and scatters candidates for the next + iteration + """ + + def __init__(self, backend) -> None: + self.backend = backend + self._target_layer_ids: Optional[List[int]] = None + self.hidden_store = SpecHiddenStore(self) + self.verifier = SpecVerifier(backend) + self.proposer = build_spec_proposer(self) + self.decode_runner = SpecDecodeRunner(self) + self.planner = self._build_decode_planner() + self._collect_dynamic_probs = self.enable_dynamic_mtp + self._dynamic_accept_stats_calls = 0 + self._full_fast_stats_count = 0 + self._full_fast_stats_interval = max( + 1, + int(os.getenv("LIGHTLLM_EAGLE3_FULL_FAST_STATS_INTERVAL", "8")), + ) + self._full_fast_accept_floor = float( + os.getenv("LIGHTLLM_EAGLE3_FULL_FAST_ACCEPT_FLOOR", "0.75") + ) + + @property + def spec_config(self) -> SpeculativeConfig: + return self.backend.spec_config + + @property + def enable_dynamic_mtp(self) -> bool: + return self.spec_config.dynamic_verify + + @property + def collect_dynamic_probs(self) -> bool: + return self._collect_dynamic_probs + + def configure_dynamic_prob_collection(self, *, plan: SpecDecodePlan) -> None: + # Eagle3's selected-token probabilities are only used to rank rows for + # confidence compaction. A profitable full-width plan neither ranks + # nor compacts rows, so use Static's cheaper argmax-only proposer. An + # ``observe`` full-width plan still records probabilities to prepare a + # possible transition back to confidence scheduling. + self._collect_dynamic_probs = self.enable_dynamic_mtp and not ( + self.planner.planner_mode == "eagle3" and plan.selection_mode == "full" + ) + + @property + def needs_intermediate_target_hidden(self) -> bool: + return self.spec_config.needs_target_layer_hidden + + def is_draft_model(self, model) -> bool: + return any(model is draft_model for draft_model in self.backend.draft_models) + + def is_block_draft_model(self, model) -> bool: + return bool(getattr(self.spec_config, "uses_block_draft_model", False)) and self.is_draft_model(model) + + def is_draft_forward(self, infer_state, model=None) -> bool: + return ( + infer_state.mtp_draft_input_hiddens is not None + or getattr(infer_state, "is_draft_model", False) + or (model is not None and self.is_draft_model(model)) + ) + + def get_capture_layer_ids(self, model, infer_state) -> List[int]: + if self.is_draft_forward(infer_state, model=model): + if self.spec_config.uses_chained_draft_models or self.spec_config.uses_recurrent_draft_model: + return [model.config["n_layer"] - 1] + return [] + if self.needs_intermediate_target_hidden: + return self._get_target_layer_ids(model) + return [] + + def create_forward_context(self, model, infer_state) -> SpecForwardContext: + return SpecForwardContext(runtime=self, model=model, infer_state=infer_state) + + def capture_hidden( + self, + *, + infer_state, + hidden: torch.Tensor, + final_hidden: torch.Tensor, + ) -> torch.Tensor: + return self.hidden_store.capture_hidden( + infer_state=infer_state, + hidden=hidden, + final_hidden=final_hidden, + ) + + def get_hidden(self, microbatch_index: int = 0) -> torch.Tensor: + return self.hidden_store.get_hidden(microbatch_index) + + def unpad_hidden(self, *, token_num: int, microbatch_index: int = 0) -> None: + self.hidden_store.unpad_hidden(token_num=token_num, microbatch_index=microbatch_index) + return + + def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: + """Allocate speculative draft-owned temporary KV slots.""" + + token_count = int(token_count) + assert token_count >= 0 + if token_count == 0: + return torch.empty((0,), dtype=torch.int32, device="cpu") + + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + if g_infer_context.radix_cache is not None: + g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) + return g_infer_context.req_manager.mem_manager.alloc(token_count) + + def build_padded_next_token_ids( + self, + *, + token_ids: Optional[torch.Tensor], + batch_size: int, + copy_len: int = None, + source_start: int = 0, + device=None, + ) -> torch.Tensor: + """Build a padded draft-token input buffer for padded DP batches.""" + + batch_size = int(batch_size) + source_start = int(source_start) + assert batch_size >= 0 + assert source_start >= 0 + if token_ids is None: + assert copy_len is None or int(copy_len) == 0 + assert device is not None + copy_len = 0 + else: + copy_len = token_ids.shape[0] - source_start if copy_len is None else int(copy_len) + assert copy_len >= 0 + assert source_start + copy_len <= token_ids.shape[0] + if device is None: + device = token_ids.device + assert copy_len <= batch_size + + padded_token_ids = torch.zeros((batch_size,), dtype=torch.int64, device=device) + if copy_len > 0: + padded_token_ids[:copy_len].copy_( + token_ids[source_start : source_start + copy_len], + non_blocking=True, + ) + return padded_token_ids + + def build_padded_eagle_step_mem_indexes( + self, + *, + eagle_mem_indexes: torch.Tensor, + step: int, + real_req_num: int, + padded_req_num: int, + ) -> torch.Tensor: + """Build one padded Eagle scratch-index column for DP decode.""" + + step = int(step) + real_req_num = int(real_req_num) + padded_req_num = int(padded_req_num) + assert step >= 0 + assert real_req_num >= 0 + assert padded_req_num >= 0 + + start = step * real_req_num + end = start + real_req_num + assert end <= eagle_mem_indexes.shape[0] + step_mem_indexes = eagle_mem_indexes[start:end] + if padded_req_num == 0: + return step_mem_indexes + + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + hold_token_memindex = g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX + hold_mem_indexes = torch.full( + (padded_req_num,), + int(hold_token_memindex), + dtype=eagle_mem_indexes.dtype, + device=eagle_mem_indexes.device, + ) + return torch.cat([step_mem_indexes, hold_mem_indexes], dim=0) + + def append_padded_eagle_step_mem_indexes( + self, + *, + model_input: ModelInput, + eagle_mem_indexes: torch.Tensor, + step: int, + real_req_num: int, + padded_req_num: int, + mtp_step: int, + ) -> torch.Tensor: + """Roll one padded Eagle scratch-index column into ModelInput.""" + + mtp_step = int(mtp_step) + assert mtp_step >= 0 + step_mem_indexes = self.build_padded_eagle_step_mem_indexes( + eagle_mem_indexes=eagle_mem_indexes, + step=step, + real_req_num=real_req_num, + padded_req_num=padded_req_num, + ) + grouped_mem_indexes = model_input.mem_indexes.view(-1, mtp_step + 1) + assert grouped_mem_indexes.shape[0] == step_mem_indexes.shape[0] + model_input.mem_indexes = torch.cat( + [grouped_mem_indexes[:, 1:], step_mem_indexes.view(-1, 1)], + dim=1, + ).view(-1) + return model_input.mem_indexes + + def prepare_draft_prefill_input( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: Optional[torch.Tensor] = None, + microbatch_index: int = 0, + ) -> ModelInput: + """Build draft prefill input from target prefill input. + + `next_token_ids`: [run_req_num] + `mtp_draft_input_hiddens`: captured target feature. The first + dimension matches the target prefill token layout after padding/unpad + handling; the second dimension is either hidden_size or + hidden_size * len(target_layer_ids). + """ + + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + return prepare_mtp_prefill_inputs( + model_input=model_input, + b_next_token_ids=next_token_ids, + mtp_draft_input_hiddens=( + self.get_hidden(microbatch_index) if mtp_draft_input_hiddens is None else mtp_draft_input_hiddens + ), + ) + + def prepare_draft_decode_input( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: Optional[torch.Tensor] = None, + microbatch_index: int = 0, + ) -> ModelInput: + """Mutate a decode ModelInput for one draft forward. + + `next_token_ids`: [verify_batch] + `mtp_draft_input_hiddens`: [verify_batch, hidden_dim_for_draft] + """ + + model_input.input_ids = next_token_ids + if mtp_draft_input_hiddens is None: + mtp_draft_input_hiddens = self.get_hidden(microbatch_index) + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + return model_input + + def graph_cache_key(self, model_context, model=None): + model = getattr(model_context, "model", None) or model + infer_state = getattr(model_context, "infer_state", model_context) + role = "draft" if self.is_draft_forward(infer_state, model=model) else "main" + disable_mtp_decode_att = bool(getattr(infer_state, "disable_mtp_decode_att", False)) + use_static_mtp_layout = bool( + role == "main" + and not disable_mtp_decode_att + and getattr(infer_state, "use_static_mtp_layout", False) + ) + return ("spec", self.spec_config.mode, role, disable_mtp_decode_att, use_static_mtp_layout) + + def get_decode_graph_mtp_step(self, model) -> int: + return self.spec_config.get_decode_graph_mtp_step( + model_config=model.config, + is_draft_model=any(model is draft_model for draft_model in self.backend.draft_models), + ) + + def get_decode_graph_warmup_mtp_step(self, model) -> int: + return self.spec_config.get_decode_graph_warmup_mtp_step( + model_config=model.config, + is_draft_model=any(model is draft_model for draft_model in self.backend.draft_models), + ) + + def export_graph_capture(self): + return self.hidden_store.export_graph_capture() + + def restore_graph_capture(self, captured_hiddens) -> None: + self.hidden_store.restore_graph_capture(captured_hiddens) + return + + def build_initial_draft_state( + self, + *, + model_input: ModelInput, + next_token_ids: torch.Tensor, + ) -> None: + self.proposer.build_initial_draft_state(model_input=model_input, next_token_ids=next_token_ids) + return + + def build_initial_draft_state_overlap( + self, + *, + model_input0: ModelInput, + next_token_ids0: torch.Tensor, + model_input1: ModelInput, + next_token_ids1: torch.Tensor, + ) -> None: + self.proposer.build_initial_draft_state_overlap( + model_input0=model_input0, + next_token_ids0=next_token_ids0, + model_input1=model_input1, + next_token_ids1=next_token_ids1, + ) + return + + def plan_decode(self, *, model_input: ModelInput, req_num: int) -> SpecDecodePlan: + """Return the static or dynamic MTP plan for one decode iteration.""" + + return self.planner.plan(req_num=req_num, original_batch_size=model_input.batch_size) + + def run_decode_speculative_forward( + self, + *, + model_input: ModelInput, + model_output: ModelOutput, + run_reqs: List, + req_num: int, + plan: SpecDecodePlan, + selected_run_reqs_cpu: Optional[torch.Tensor], + next_token_ids: torch.Tensor, + next_token_logprobs: torch.Tensor, + next_token_ranks: torch.Tensor, + copy_next_token_infos: Callable[ + [torch.Tensor, torch.Tensor, torch.Tensor], + Tuple[torch.Tensor, torch.Tensor, torch.Tensor], + ], + ) -> SpecDecodeForwardState: + return self.decode_runner.run_speculative_forward( + model_input=model_input, + model_output=model_output, + run_reqs=run_reqs, + req_num=req_num, + plan=plan, + selected_run_reqs_cpu=selected_run_reqs_cpu, + next_token_ids=next_token_ids, + next_token_logprobs=next_token_logprobs, + next_token_ranks=next_token_ranks, + copy_next_token_infos=copy_next_token_infos, + ) + + def resolve_decode_pre_post_reqs(self, *, state: SpecDecodeForwardState, decode_reqs: List): + return self.decode_runner.resolve_pre_post_reqs(state=state, decode_reqs=decode_reqs) + + def finish_decode_post( + self, + *, + state: SpecDecodeForwardState, + req_num: int, + run_reqs: List, + ) -> SpecDecodePostState: + return self.decode_runner.finish_post(state=state, req_num=req_num, run_reqs=run_reqs) + + def prepare_decode_model_input( + self, + *, + model_input: ModelInput, + req_num: int, + plan: SpecDecodePlan, + ): + """Apply dynamic MTP row compaction when the planner selects it.""" + + if not plan.is_dynamic: + return model_input, None + + model_input.use_static_mtp_layout = False + if plan.selection_mode in {"full", "observe"}: + assert plan.dynamic_batch_size == model_input.batch_size + assert plan.pre_draft_step == self.backend.mtp_step + model_input.use_static_mtp_layout = True + return model_input, None + + self._clear_stale_dynamic_token_probs(pre_draft_step=plan.pre_draft_step) + + from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_mtp_model_input + + selection_verify_step = plan.pre_draft_step + if plan.selection_mode == "prefix": + assert plan.dynamic_batch_size % req_num == 0 + selection_verify_step = plan.dynamic_batch_size // req_num - 1 + + model_input, selected_run_reqs = prepare_dynamic_mtp_model_input( + model_input=model_input, + req_num=req_num, + dynamic_batch_size=plan.dynamic_batch_size, + req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + req_to_next_token_probs=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs, + verify_step=selection_verify_step, + use_prefix_selection=plan.selection_mode == "prefix", + ) + return model_input, selected_run_reqs + + def async_copy_selected_run_reqs(self, selected_run_reqs: Optional[torch.Tensor]): + if selected_run_reqs is None: + return None + from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + + return g_pin_mem_manager.async_copy_from_gpu_tensor( + key="selected_run_reqs", + gpu_tensor=selected_run_reqs, + ) + + def build_decode_req_lists( + self, + *, + original_run_reqs, + selected_run_reqs_cpu: Optional[torch.Tensor], + accepted_index_cpu: torch.Tensor, + ): + """Build post-handle request lists after optional dynamic MTP compaction.""" + + if self.enable_dynamic_mtp and selected_run_reqs_cpu is not None: + assert selected_run_reqs_cpu is not None + selected_run_reqs_cpu_numpy = selected_run_reqs_cpu.numpy() + run_reqs = [ + original_run_reqs[i] for i in range(len(original_run_reqs)) if selected_run_reqs_cpu_numpy[i] == 1 + ] + else: + run_reqs = original_run_reqs + + accepted_index_cpu_numpy = accepted_index_cpu.numpy() + verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu_numpy[i] == 1] + return run_reqs, verify_ok_reqs + + def build_decode_free_mem_indexes_cpu( + self, + *, + model_input: ModelInput, + selected_run_reqs_cpu: Optional[torch.Tensor], + accepted_index_cpu: torch.Tensor, + ) -> torch.Tensor: + mem_indexes_cpu = model_input.mem_indexes_cpu + if not self.enable_dynamic_mtp or selected_run_reqs_cpu is None: + return mem_indexes_cpu[accepted_index_cpu == 0] + + assert selected_run_reqs_cpu is not None + selected_mask = selected_run_reqs_cpu.to(dtype=torch.bool) + accepted_mask = accepted_index_cpu.to(dtype=torch.bool) + selected_mem_indexes_cpu = mem_indexes_cpu[selected_mask] + assert selected_mem_indexes_cpu.shape[0] == accepted_mask.shape[0] + + unselected_mem_indexes_cpu = mem_indexes_cpu[~selected_mask] + rejected_selected_mem_indexes_cpu = selected_mem_indexes_cpu[~accepted_mask] + if len(unselected_mem_indexes_cpu) == 0: + return rejected_selected_mem_indexes_cpu + if len(rejected_selected_mem_indexes_cpu) == 0: + return unselected_mem_indexes_cpu + return torch.cat([unselected_mem_indexes_cpu, rejected_selected_mem_indexes_cpu], dim=0) + + def update_dynamic_accept_stats( + self, + *, + req_num: int, + run_reqs, + accepted_index_cpu: torch.Tensor, + mtp_accept_len_cpu: torch.Tensor, + dynamic_batch_size: Optional[int], + verify_step: Optional[int] = None, + selection_mode: str = "confidence", + ) -> None: + if not self.enable_dynamic_mtp: + return + + assert dynamic_batch_size is not None + assert len(run_reqs) == accepted_index_cpu.shape[0] + assert mtp_accept_len_cpu.shape[0] == req_num + accept_lengths = mtp_accept_len_cpu.numpy() + accept_count = int(accept_lengths.sum()) + total_count = int(dynamic_batch_size) + is_full_verify = dynamic_batch_size == req_num * (self.backend.mtp_step + 1) + + # The first target decode has no preceding draft proposal, so every + # request structurally accepts only its base token. Treating that + # cold-start iteration as a K-wide acceptance sample makes a highly + # predictable workload look maximally hard and can collapse the + # controller before the first real proposal is verified. + self._dynamic_accept_stats_calls += 1 + if self._dynamic_accept_stats_calls == 1: + return + + # In the profitable full-width state, the per-iteration trace already + # records exact acceptance. Planner EMAs only need periodic refreshes + # while acceptance remains high. A workload drop bypasses sampling + # immediately, so contraction still reacts on the first bad batch. + if self.planner.planner_mode == "eagle3" and selection_mode == "full": + self._full_fast_stats_count += 1 + current_accept_ratio = accept_count / max(1, total_count) + if ( + current_accept_ratio >= self._full_fast_accept_floor + and self._full_fast_stats_count % self._full_fast_stats_interval != 0 + ): + return + + self.planner.update_req_num_to_dynamic_batch_size_to_accept_ratio( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + accept_ratio=accept_count / total_count, + **({"verify_step": verify_step} if self.planner.planner_mode == "eagle3" else {}), + ) + + update_full_verify_tokens_per_req = getattr( + self.planner, + "update_full_verify_tokens_per_req", + None, + ) + if update_full_verify_tokens_per_req is not None and is_full_verify: + update_full_verify_tokens_per_req( + accept_count / req_num, + req_num=req_num, + ) + + update_observed_iteration_stats = getattr( + self.planner, + "update_observed_iteration_stats", + None, + ) + if update_observed_iteration_stats is not None: + update_observed_iteration_stats( + tokens_per_req=accept_count / req_num, + verify_rows_per_req=dynamic_batch_size / req_num, + is_full_verify=is_full_verify, + req_num=req_num, + ) + + # Eagle3 uses these values as an unbiased survival curve. Updating it + # from confidence-selected dynamic rows would bias every depth upward; + # full-width warmup/probe iterations are the valid samples. + update_verified_batch_prefix_stats = getattr( + self.planner, + "update_verified_batch_prefix_stats", + None, + ) + if is_full_verify and update_verified_batch_prefix_stats is not None: + update_verified_batch_prefix_stats( + verify_and_accept_lengths=[ + (self.backend.mtp_step + 1, int(accept_len)) for accept_len in accept_lengths + ], + ) + elif self.planner.planner_mode != "eagle3": + verify_rows_per_req = max(1, int(round(dynamic_batch_size / req_num))) + for accept_len in accept_lengths: + self.planner.update_verified_prefix_stats( + verify_len=verify_rows_per_req, + accept_len=int(accept_len), + ) + return + + def needs_schedule_probs_cpu(self) -> bool: + """Whether this planner consumes proposal confidence on the CPU.""" + + planner_needs_schedule_probs_cpu = getattr(self.planner, "needs_schedule_probs_cpu", None) + if planner_needs_schedule_probs_cpu is not None: + return bool(planner_needs_schedule_probs_cpu()) + return callable(getattr(self.planner, "update_predicted_schedule_probs", None)) + + def update_dynamic_schedule_stats( + self, + *, + req_num: int, + schedule_probs_cpu: Optional[torch.Tensor], + ) -> None: + if not self.enable_dynamic_mtp or schedule_probs_cpu is None: + return + + update_predicted_schedule_probs = getattr(self.planner, "update_predicted_schedule_probs", None) + if update_predicted_schedule_probs is None: + return + + update_predicted_schedule_probs( + schedule_probs=schedule_probs_cpu, + req_num=req_num, + ) + return + + def propose_next( + self, + *, + main_model_input: ModelInput, + main_model_output: Optional[ModelOutput] = None, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + verify_result: Optional[SpecVerifyResult] = None, + ) -> SpecProposal: + return self.proposer.propose_next( + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=draft_step, + verify_result=verify_result, + ) + + def verify_target_tokens( + self, + *, + new_next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + ) -> SpecVerifyResult: + return self.verifier.verify_target_tokens( + new_next_token_ids=new_next_token_ids, + b_req_idx=b_req_idx, + b_req_mtp_start_loc=b_req_mtp_start_loc, + ) + + def build_all_next_token_probs( + self, + *, + next_token_logprobs: torch.Tensor, + proposal: SpecProposal, + draft_step: int, + ) -> Optional[torch.Tensor]: + """Build selected-token probability matrix for dynamic MTP scatter. + + Output shape is [verify_batch, mtp_step + 1]. Column 0 is the target + token probability, fixed to 1 because the target sample is always the + base accepted position. Draft columns store selected-token + probabilities from each proposer step. + """ + + if not self.enable_dynamic_mtp or not self.collect_dynamic_probs: + return None + + schedule_probs = proposal.schedule_probs if proposal.schedule_probs is not None else proposal.draft_probs + assert schedule_probs is not None + + all_next_token_probs = torch.zeros( + size=(next_token_logprobs.shape[0], self.backend.mtp_step + 1), + dtype=torch.float32, + device=next_token_logprobs.device, + ) + all_next_token_probs[:, 0] = 1.0 + + if isinstance(schedule_probs, torch.Tensor): + assert schedule_probs.shape == (next_token_logprobs.shape[0], draft_step) + if draft_step > 0: + all_next_token_probs[:, 1 : draft_step + 1] = schedule_probs + return all_next_token_probs + + assert len(schedule_probs) == draft_step + for step_idx, step_probs in enumerate(schedule_probs): + all_next_token_probs[:, step_idx + 1] = step_probs + return all_next_token_probs + + def pad_all_next_token_ids(self, *, token_ids: torch.Tensor, draft_step: int) -> torch.Tensor: + """Pad dynamic proposal ids to the static MTP width before scatter.""" + + if not self.enable_dynamic_mtp or draft_step >= self.backend.mtp_step: + return token_ids + + append_next_token_ids = torch.ones( + size=(token_ids.shape[0], self.backend.mtp_step - draft_step), + dtype=token_ids.dtype, + device=token_ids.device, + ) + return torch.cat([token_ids, append_next_token_ids], dim=-1) + + def scatter_next_tokens( + self, + *, + b_req_mtp_start_loc: torch.Tensor, + all_next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + mtp_accept_len: torch.Tensor, + all_next_token_probs: Optional[torch.Tensor] = None, + ) -> None: + self.verifier.scatter_next_tokens( + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=b_req_idx, + mtp_accept_len=mtp_accept_len, + all_next_token_probs=all_next_token_probs, + ) + return + + def scatter_token_id_steps( + self, + *, + token_id_steps: List[torch.Tensor], + b_req_mtp_start_loc: torch.Tensor, + b_req_idx: torch.Tensor, + mtp_accept_len: torch.Tensor, + row_count: int = None, + ) -> torch.Tensor: + """Stack proposal token columns and scatter them for the next verify.""" + + all_next_token_ids = torch.stack(token_id_steps, dim=1) + if row_count is not None: + all_next_token_ids = all_next_token_ids[: int(row_count), :] + self.scatter_next_tokens( + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=b_req_idx, + mtp_accept_len=mtp_accept_len, + ) + return all_next_token_ids + + def _get_target_layer_ids(self, model) -> List[int]: + if self._target_layer_ids is not None: + return self._target_layer_ids + + draft_model_dir = self.backend.args.mtp_draft_model_dir + if isinstance(draft_model_dir, list): + draft_model_dir = draft_model_dir[0] + + if draft_model_dir: + with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: + draft_config = json.load(json_file) + target_layer_ids = draft_config.get("target_layer_ids") + if target_layer_ids is not None: + self._target_layer_ids = [int(layer_id) for layer_id in target_layer_ids] + return self._target_layer_ids + + self._target_layer_ids = [1, model.config["n_layer"] // 2 - 1, model.config["n_layer"] - 4] + return self._target_layer_ids + + def _build_decode_planner(self): + if self.enable_dynamic_mtp: + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + assert g_infer_context.dynamic_mtp_planner is not None + return g_infer_context.dynamic_mtp_planner + return FixedMTPPlanner(self.backend.mtp_step) + + def _clear_stale_dynamic_token_probs(self, *, pre_draft_step: int) -> None: + # Columns after the previous draft length are stale and must not be + # sampled by dynamic row compaction in the current target forward. + self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs[ + :, (pre_draft_step + 1) : + ].fill_(0.0) + return + + +def build_spec_runtime(backend) -> SpecRuntime: + return SpecRuntime(backend) diff --git a/lightllm/server/router/model_infer/speculative/state.py b/lightllm/server/router/model_infer/speculative/state.py new file mode 100644 index 0000000000..af16ce8510 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/state.py @@ -0,0 +1,130 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, List, Optional + +import torch + +from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.speculative.runtime import SpecRuntime + + +class SpecHiddenStore: + """Stores target-model features that are passed into draft-model forwards. + + LightLLM's service path does not materialize HuggingFace-style + `hidden_states`. Instead, BaseModel calls SpecForwardContext while the + target model is running. + + Captured tensors are keyed by microbatch index: + - vanilla MTP consumes the final target hidden state: + [token_num, hidden_size] + - Eagle3 / DSpark-style draft models consume selected target layers + concatenated on the hidden dimension after TP/SP all-gather: + [token_num, hidden_size * len(target_layer_ids)] + + Draft-model forwards are identified by `mtp_draft_input_hiddens is not + None`; their final hidden can also be captured because chained MTP drafts + pass one draft's hidden into the next draft model. + """ + + def __init__(self, runtime: "SpecRuntime") -> None: + self.runtime = runtime + self._captured_hiddens = {} + + def select_hidden(self, *, infer_state, hidden: torch.Tensor, final_hidden: torch.Tensor) -> torch.Tensor: + if self.runtime.is_draft_forward(infer_state): + return final_hidden + if self.runtime.needs_intermediate_target_hidden: + return hidden + return final_hidden + + def capture_hidden(self, *, infer_state, hidden: torch.Tensor, final_hidden: torch.Tensor) -> torch.Tensor: + selected_hidden = self.select_hidden( + infer_state=infer_state, + hidden=hidden, + final_hidden=final_hidden, + ) + self._captured_hiddens[infer_state.microbatch_index] = selected_hidden + return selected_hidden + + def get_hidden(self, microbatch_index: int = 0) -> torch.Tensor: + hidden = self._captured_hiddens.get(microbatch_index) + assert hidden is not None + return hidden + + def unpad_hidden(self, *, token_num: int, microbatch_index: int = 0) -> None: + hidden = self._captured_hiddens.get(microbatch_index) + if hidden is not None and hidden.shape[0] > token_num: + self._captured_hiddens[microbatch_index] = hidden[0:token_num] + return + + def export_graph_capture(self): + captured_hiddens = { + microbatch_index: tensor_to_no_ref_tensor(hidden) + for microbatch_index, hidden in self._captured_hiddens.items() + } + self._captured_hiddens = dict(captured_hiddens) + return captured_hiddens + + def restore_graph_capture(self, captured_hiddens) -> None: + if captured_hiddens is not None: + self._captured_hiddens = dict(captured_hiddens) + return + + +class SpecForwardContext: + """Per-forward hidden capture context used by BaseModel. + + The target model and the draft model exchange only tensors, not model + outputs. During target forward, BaseModel calls `add_hidden` after each + transformer layer. The runtime decides which layer ids matter for the + active draft algorithm. At the end of forward, BaseModel calls `capture` + with: + - `hidden`: selected intermediate target feature, shape + [token_num, hidden_size * selected_layer_num] after all-gather + - `final_hidden`: final target feature, shape [token_num, hidden_size] + + Vanilla MTP uses `final_hidden`; Eagle3/DSpark-style proposers use + `hidden`. + """ + + def __init__(self, *, runtime: "SpecRuntime", model, infer_state) -> None: + self.runtime = runtime + self.model = model + self.infer_state = infer_state + self.layer_ids = runtime.get_capture_layer_ids(model, infer_state) + self.layer_hiddens: List[torch.Tensor] = [] + + def add_hidden(self, *, layer_index: int, layer_num: int, hidden: torch.Tensor) -> None: + if layer_index not in self.layer_ids: + return + + if layer_index == layer_num - 1: + self.layer_hiddens.append(hidden) + else: + self.layer_hiddens.append(hidden.clone()) + return + + def build_layer_hidden(self) -> Optional[torch.Tensor]: + if not self.layer_hiddens: + return None + if len(self.layer_hiddens) == 1: + return self.layer_hiddens[0] + return torch.cat(self.layer_hiddens, dim=-1) + + def capture(self, *, hidden: torch.Tensor, final_hidden: torch.Tensor) -> torch.Tensor: + return self.runtime.capture_hidden( + infer_state=self.infer_state, + hidden=hidden, + final_hidden=final_hidden, + ) + + def capture_final_hidden(self, final_hidden: torch.Tensor) -> None: + assert not self.layer_ids, ( + f"{self.runtime.spec_config.mode} needs intermediate hidden layers and does not support " + "microbatch overlap forward now" + ) + self.capture(hidden=final_hidden, final_hidden=final_hidden) + return diff --git a/lightllm/server/router/model_infer/speculative/verifier.py b/lightllm/server/router/model_infer/speculative/verifier.py new file mode 100644 index 0000000000..1a29db0ff5 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/verifier.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional + +import torch + +from lightllm.common.basemodel.triton_kernel.mtp_utils import mtp_scatter_next_token_ids, mtp_verify + + +@dataclass +class SpecVerifyResult: + """GPU verification result produced after target decode. + + `accept_len` is per logical request and includes the target token at + position 0: + + accept_len: [logical_req_num] + + `accepted_index` is per verified row in the target batch. It marks rows + whose token should be committed and post-processed: + + accepted_index: [verify_batch] + """ + + accept_len: torch.Tensor + accepted_index: torch.Tensor + + +class SpecVerifier: + """Service verifier for LightLLM MTP layout. + + DeepSpec verifies by running target logits over + [current_token + draft_tokens] and applying rejection sampling. LightLLM's + service path stores pending candidates in req_sampling_params_manager and + uses Triton kernels to: + - compare the target sampled tokens with previously scattered candidates + - compute per-request accepted prefix length + - scatter the newly proposed candidates for the next decode iteration + """ + + def __init__(self, backend) -> None: + self.backend = backend + + def verify_target_tokens( + self, + *, + new_next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + ) -> SpecVerifyResult: + """Verify target sampled ids against previously proposed ids. + + Inputs: + - `new_next_token_ids`: target sampled ids, shape [verify_batch] + - `b_req_idx`: request ids for each target row, shape [verify_batch] + - `b_req_mtp_start_loc`: first row of each logical request in the + MTP-expanded target batch, shape [logical_req_num] + """ + + mtp_accept_len, accepted_index = mtp_verify( + req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + new_next_token_ids=new_next_token_ids, + b_req_idx=b_req_idx, + ) + return SpecVerifyResult(accept_len=mtp_accept_len, accepted_index=accepted_index) + + def scatter_next_tokens( + self, + *, + b_req_mtp_start_loc: torch.Tensor, + all_next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + mtp_accept_len: torch.Tensor, + all_next_token_probs: Optional[torch.Tensor] = None, + ) -> None: + """Scatter target+draft candidates into per-request next-token buffers. + + Inputs: + - `all_next_token_ids`: [verify_batch, mtp_step + 1]. Column 0 is the + target sampled token from this iteration; remaining columns are draft + candidates padded to the static MTP width when dynamic MTP produces a + shorter proposal. + - `all_next_token_probs`: optional [verify_batch, mtp_step + 1]. + Current dynamic MTP stores selected-token probabilities, not full + vocab distributions. + """ + + mtp_scatter_next_token_ids( + req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=b_req_idx, + mtp_accept_len=mtp_accept_len, + # A profitable full-width Eagle3 plan deliberately uses Static's + # argmax-only proposer and therefore has no confidence matrix to + # scatter. Leave the probability buffer untouched in that mode; + # the next observe/confidence plan refreshes it before compaction. + req_to_next_token_probs=( + self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs + if all_next_token_probs is not None + else None + ), + all_next_token_probs=all_next_token_probs, + ) + return diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 33692f431e..b16e1612ac 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -5,6 +5,7 @@ from easydict import EasyDict from functools import lru_cache from lightllm.utils.log_utils import init_logger +from lightllm.common.speculative import SpeculativeConfig logger = init_logger(__name__) @@ -221,6 +222,26 @@ def enable_diverse_mode_gqa_decode_fast_kernel() -> bool: return get_env_start_args().diverse_mode and "int8kv" == get_env_start_args().llm_kv_type +@lru_cache(maxsize=None) +def enable_dynamic_mtp_verify() -> bool: + """ + 启用动态 MTP 长度验证功能 + 在 MTP 模式下,根据每步的 prob 分布动态调整验证长度 + 通过启动参数 --mtp_dynamic_verify 控制;DSpark 模式固定使用 + confidence-scheduled dynamic verify。 + """ + return SpeculativeConfig.from_args(get_env_start_args()).dynamic_verify + + +@lru_cache(maxsize=None) +def enable_triton_mtp_kernel() -> bool: + """ + 启用 Triton MTP 解码专用 kernel + 通过启动参数 --mtp_step > 0 和 --llm_decode_att_backend=triton 控制 + """ + return (get_env_start_args().mtp_step > 0) and ("triton" in get_env_start_args().llm_decode_att_backend) + + @lru_cache(maxsize=None) def get_disk_cache_prompt_limit_length(): return int(os.getenv("LIGHTLLM_DISK_CACHE_PROMPT_LIMIT_LENGTH", 2048)) @@ -246,14 +267,20 @@ def enable_cpu_cache_numa_interleave() -> bool: @lru_cache(maxsize=None) def get_added_mtp_kv_layer_num() -> int: - # mtp 模式下需要在mem manger上扩展draft model使用的layer - added_mtp_layer_num = 0 - if get_env_start_args().mtp_mode == "eagle_with_att": - added_mtp_layer_num += 1 - elif get_env_start_args().mtp_mode == "vanilla_with_att": - added_mtp_layer_num += get_env_start_args().mtp_step - - return added_mtp_layer_num + args = get_env_start_args() + spec_config = SpeculativeConfig.from_args(args) + if not spec_config.uses_attention_draft: + return 0 + if spec_config.is_dflash or spec_config.is_dspark or spec_config.is_eagle3: + draft_model_dir = args.mtp_draft_model_dir + if isinstance(draft_model_dir, list): + draft_model_dir = draft_model_dir[0] + if not draft_model_dir: + return spec_config.draft_model_count + with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: + draft_config = json.load(json_file) + return int(draft_config.get("num_hidden_layers", draft_config.get("n_layer", spec_config.draft_model_count))) + return spec_config.draft_model_count @lru_cache(maxsize=None) diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index e81caafe7a..0b330b7602 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -16,17 +16,17 @@ get_added_mtp_kv_layer_num, ) from lightllm.utils.log_utils import init_logger +from lightllm.common.speculative import SpeculativeConfig from lightllm.utils.config_utils import get_num_key_value_heads, get_head_dim, get_layer_num, is_linear_att_mixed_model from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class from lightllm.common.kv_cache_mem_manager import ( MemoryManager, PPLINT8KVMemoryManager, - PPLINT4KVMemoryManager, Deepseek2MemoryManager, Qwen3NextMemManager, ) -from typing import List, Tuple, Optional +from typing import List, Tuple from tqdm import tqdm from lightllm.utils.auto_shm_cleanup import register_sysv_shm_for_cleanup from lightllm.utils.dist_utils import get_current_device_id @@ -119,7 +119,8 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": logger.error(f"not support mem manager: {mem_manager_class} for cpu kv cache") raise Exception(f"not support mem manager: {mem_manager_class} for cpu kv cache") - if args.mtp_mode is not None: + spec_config = SpeculativeConfig.from_args(args) + if spec_config.enabled: # TODO 可能会存在不同mtp模式的精度问题 if not is_linear_att_mixed_model(args.model_dir): # 对于非 linear att 混合模型,需要额外增加 mtp 的 kv 层数, diff --git a/unit_tests/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/test_mtp_diverse.py b/unit_tests/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/test_mtp_diverse.py new file mode 100644 index 0000000000..b25088790c --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/att/decode_att/gqa/mtp_diverse/test_mtp_diverse.py @@ -0,0 +1,195 @@ +""" +MTP Diverse Attention Unit Test + +测试 MTP diverse attention 算子的正确性和性能。 + +Single Token Mode: +- 每个请求只有 1 个 Q token +- 组内第 i 个请求只能看到前 i 个 KV +- b_mark_shared_group 标记: + - 0: 组内非最后一个请求(跳过) + - N>=1: 一个 N 人组的最后一个请求(需要计算) +""" +import pytest +import torch + + +def gqa_attention_reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + req_to_tokens: torch.Tensor, + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, +) -> torch.Tensor: + """ + Reference GQA attention implementation for single token mode + + q: [batch_size, num_heads, head_dim] - 每个请求 1 个 Q token + k, v: [kv_pool_size, kv_head_num, head_dim] - KV 池 + req_to_tokens: [batch, max_kv_len] - 每个请求的 KV 索引 + b_seq_len: [batch] - 每个请求可见的 KV 数量 + """ + batch_size = b_seq_len.shape[0] + num_heads = q.shape[1] + kv_head_num = k.shape[1] + head_dim = q.shape[2] + gqa_group_size = num_heads // kv_head_num + + # 输出:[batch, num_heads, head_dim] + output = torch.zeros((batch_size, num_heads, head_dim), dtype=q.dtype, device=q.device) + + for b in range(batch_size): + seq_len = b_seq_len[b].item() # 这个请求可见的 KV 数量 + + # 获取这个请求的 KV 索引:[seq_len] + kv_indices = req_to_tokens[b, :seq_len] + + # 获取 K 和 V:[seq_len, kv_head_num, head_dim] + k_batch = k[kv_indices] + v_batch = v[kv_indices] + + # Process each Q head + for h in range(num_heads): + kv_h = h // gqa_group_size + k_h = k_batch[:, kv_h, :] # [seq_len, head_dim] + v_h = v_batch[:, kv_h, :] # [seq_len, head_dim] + + q_h = q[b, h] # [head_dim] + + # Compute attention scores + att = torch.matmul(q_h, k_h.transpose(0, 1)) # [seq_len] + att = att / (head_dim ** 0.5) + att = torch.softmax(att, dim=-1) + out_h = torch.matmul(att, v_h) # [head_dim] + + output[b, h] = out_h + + return output + + +def setup_mtp_test_data(kv_len, group_size, batch_groups, test_dtype=torch.bfloat16, device="cuda", seed=42): + """ + 设置 MTP 测试数据 - Single Token Mode + + Single Token Mode 的特点: + - 每个请求有 1 个 Q token + - 组内第 i 个请求只能看到前 i+1 个 KV + - 组内请求共享相同的 KV 前缀 + + b_mark_shared_group 标记: + - 0: 组内非最后一个请求(跳过计算) + - N>=1: 一个 N 人组的最后一个请求(需要计算) + """ + torch.manual_seed(seed) + + num_heads = 32 + kv_head_num = 4 # gqa_group_size = 8 + head_dim = 128 + + batch_size = batch_groups * group_size + + # KV 池:[batch_groups * kv_len, kv_head_num, head_dim] + # 每个组使用不同的 KV 范围 + kv_pool_size = batch_groups * kv_len + k = torch.randn(size=(kv_pool_size, kv_head_num, head_dim), dtype=test_dtype, device=device) + v = torch.randn(size=(kv_pool_size, kv_head_num, head_dim), dtype=test_dtype, device=device) + + # req_to_tokens: [batch, max_kv_len] - 每个请求的 KV 索引 + max_kv_len = group_size # 每个请求最多有 group_size 个 KV(可见范围) + req_to_tokens = torch.zeros((batch_size, max_kv_len), dtype=torch.int32, device=device) + + b_req_idx = torch.arange(batch_size, dtype=torch.int32, device=device) + b_seq_len = torch.zeros(batch_size, dtype=torch.int32, device=device) # 每个请求可见的 KV 数量 + b_mark_shared_group = torch.zeros(batch_size, dtype=torch.int32, device=device) + + # Q: [batch_size, num_heads, head_dim] - 每个请求 1 个 Q token + q = torch.randn(size=(batch_size, num_heads, head_dim), dtype=test_dtype, device=device) + + for group_idx in range(batch_groups): + group_start = group_idx * group_size + kv_base = group_idx * kv_len + + for member_idx in range(group_size): + batch_idx = group_start + member_idx + # 第 member_idx 个请求可以看到前 member_idx+1 个 KV + b_seq_len[batch_idx] = member_idx + 1 + + # 设置 KV 索引 - 组内请求共享相同的 KV 前缀 + for kv_pos in range(member_idx + 1): + req_to_tokens[batch_idx, kv_pos] = kv_base + kv_pos + + # b_mark_shared_group: + # - 0: 组内非最后一个请求 + # - group_size: 组内最后一个请求 + if member_idx == group_size - 1: + b_mark_shared_group[batch_idx] = group_size + else: + b_mark_shared_group[batch_idx] = 0 + + return q, k, v, req_to_tokens, b_req_idx, b_seq_len, b_mark_shared_group + + +def mtp_diverse_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + req_to_tokens: torch.Tensor, + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + b_mark_shared_group: torch.Tensor, +) -> torch.Tensor: + """ + MTP Diverse Attention 调用入口 + """ + from lightllm.common.basemodel.triton_kernel.att.decode_att.gqa.mtp_diverse import ( + token_decode_attention_mtp_diverse_single_token, + ) + + return token_decode_attention_mtp_diverse_single_token( + q=q, + k=k, + v=v, + Req_to_tokens=req_to_tokens, + B_req_idx=b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + ) + + +@pytest.mark.parametrize("kv_len", [100, 256]) +@pytest.mark.parametrize("group_size", [2, 3]) +@pytest.mark.parametrize("batch_groups", [1, 2, 4]) +def test_mtp_diverse_vs_reference(kv_len, group_size, batch_groups): + """ + 测试 MTP diverse attention 与 reference 实现的正确性对比 + """ + test_dtype = torch.bfloat16 + device = "cuda" + + q, k, v, req_to_tokens, b_req_idx, b_seq_len, b_mark_shared_group = setup_mtp_test_data( + kv_len, group_size, batch_groups, test_dtype, device + ) + + # 运行 reference 实现(计算所有请求) + reference_out = gqa_attention_reference( + q=q, + k=k, + v=v, + req_to_tokens=req_to_tokens, + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + ) + + # 运行 MTP diverse(只计算组内最后一个请求) + diverse_out = mtp_diverse_attention( + q=q, + k=k, + v=v, + req_to_tokens=req_to_tokens, + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + ) + + torch.testing.assert_close(diverse_out, reference_out, atol=1e-2, rtol=1e-2) diff --git a/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py b/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py index 3e2555b339..816c807c24 100644 --- a/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py +++ b/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py @@ -54,7 +54,7 @@ def test_token_decode_attention_flash_decoding_diverse_matches_normal_decode(sha ) num_heads = 32 - kv_head_num = 8 + kv_head_num = 2 # gqa_group_size = 16,满足 Triton tl.dot 的 M >= 16 要求 mark_shared_group_size = 3 seq_len = 3547 head_dim = 128 @@ -118,6 +118,7 @@ def test_token_decode_attention_flash_decoding_diverse_matches_normal_decode(sha cache_v_scale=cache_v_scale, alloc_tensor_func=alloc_tensor_func, ) + # 运行 diverse 版本 diverse_out = diverse_attention( q=q.clone(), @@ -129,11 +130,4 @@ def test_token_decode_attention_flash_decoding_diverse_matches_normal_decode(sha alloc_tensor_func=alloc_tensor_func, ) - print(f"\nshared_seq_len={shared_seq_len}\nbatch_size={batch_size}") - print(f"normal_out: {normal_out[0, 0, :4]}") - print(f"diverse_out: {diverse_out[0, 0, :4]}") - print(f"max diff: {(normal_out - diverse_out).abs().max()}") - - assert torch.allclose( - normal_out, diverse_out, atol=1e-2, rtol=1e-2 - ), f"diverse vs normal decode mismatch for shared_seq_len={shared_seq_len}" + torch.testing.assert_close(normal_out, diverse_out, atol=1e-2, rtol=1e-2) diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py new file mode 100644 index 0000000000..e99c0ee357 --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py @@ -0,0 +1,232 @@ +import torch +import pytest +import triton +import numpy as np + +from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import ( + _fwd_kernel_cumprod_probs, + sample_dynamic_mtp_req_mask, +) + + +def _reference_cumprod_probs(req_to_next_token_probs, b_req_idx, mtp_step: int) -> torch.Tensor: + probs = req_to_next_token_probs.clone() + req_num = b_req_idx.shape[0] // (mtp_step + 1) + for req_i in range(req_num): + req_idx = int(b_req_idx[req_i * (mtp_step + 1)].item()) + row = probs[req_idx, : mtp_step + 1].clone() + row[0] = 1.0 + row[1:] = torch.clamp(row[1:], min=0.01, max=0.99) + probs[req_idx, : mtp_step + 1] = torch.cumprod(row, dim=0) + return probs + + +def _flat_cumprod_probs( + b_req_idx: torch.Tensor, + req_to_next_token_probs: torch.Tensor, + mtp_step: int, +) -> torch.Tensor: + probs = _reference_cumprod_probs(req_to_next_token_probs, b_req_idx, mtp_step) + req_num = b_req_idx.shape[0] // (mtp_step + 1) + all_num = req_num * (mtp_step + 1) + flat_probs = [] + for offset in range(all_num): + req_idx = int(b_req_idx[offset].item()) + mtp_index = offset % (mtp_step + 1) + flat_probs.append(probs[req_idx, mtp_index]) + return torch.stack(flat_probs) + + +def _assert_topk_mask(select: torch.Tensor, flat_probs: torch.Tensor, dynamic_batch_size: int) -> None: + k = dynamic_batch_size + assert int(select.sum().item()) == k + selected_scores = flat_probs[select.bool()] + unselected_scores = flat_probs[(select == 0).bool()] + if unselected_scores.numel() > 0: + assert selected_scores.min() >= unselected_scores.max() - 1e-5 + + +def _make_batch_probs(req_num: int, mtp_step: int, rows): + max_req = req_num + probs = torch.zeros((max_req + 1, 16), dtype=torch.float32, device="cuda") + for req_idx, row in enumerate(rows): + probs[req_idx, : mtp_step + 1] = torch.tensor(row, dtype=torch.float32, device="cuda") + b_req_idx = torch.arange(req_num, dtype=torch.int32, device="cuda").repeat_interleave(mtp_step + 1) + return probs, b_req_idx + + +@pytest.mark.parametrize("mtp_step", [1, 3]) +def test_cumprod_probs_kernel(mtp_step: int): + req_num = 2 + probs, b_req_idx = _make_batch_probs( + req_num, + mtp_step, + rows=[ + [1.0] + [0.5] * mtp_step, + [1.0] + [0.2] * mtp_step, + ], + ) + probs_clone = probs.clone() + _fwd_kernel_cumprod_probs[(req_num,)]( + req_to_next_token_probs=probs_clone, + req_to_next_token_probs_stride=probs_clone.stride(0), + b_req_idx=b_req_idx, + mtp_step=mtp_step, + BLOCK_SIZE=triton.next_power_of_2(mtp_step + 1), + num_warps=1, + num_stages=1, + ) + expected = _reference_cumprod_probs(probs, b_req_idx, mtp_step) + assert torch.allclose(probs_clone[:, : mtp_step + 1], expected[:, : mtp_step + 1], rtol=1e-5, atol=1e-5) + + +def test_cumprod_probs_clamps_invalid_values(): + mtp_step = 2 + req_num = 1 + probs, b_req_idx = _make_batch_probs(req_num, mtp_step, rows=[[1.0, 0.0, 1.5]]) + _fwd_kernel_cumprod_probs[(req_num,)]( + req_to_next_token_probs=probs, + req_to_next_token_probs_stride=probs.stride(0), + b_req_idx=b_req_idx, + mtp_step=mtp_step, + BLOCK_SIZE=triton.next_power_of_2(mtp_step + 1), + num_warps=1, + num_stages=1, + ) + row = probs[0, : mtp_step + 1] + # Index 0 is the always-accepted target sample; clamping applies only to drafts. + assert row[0].item() == pytest.approx(1.0) + assert row[1].item() == pytest.approx(0.01, rel=1e-4) + assert row[2].item() == pytest.approx(0.01 * 0.99, rel=1e-4) + + +def test_cumprod_probs_clamps_boundary_values(): + mtp_step = 3 + req_num = 1 + probs, b_req_idx = _make_batch_probs(req_num, mtp_step, rows=[[1.0, 0.995, 0.005, 0.5]]) + raw_probs = probs.clone() + _fwd_kernel_cumprod_probs[(req_num,)]( + req_to_next_token_probs=probs, + req_to_next_token_probs_stride=probs.stride(0), + b_req_idx=b_req_idx, + mtp_step=mtp_step, + BLOCK_SIZE=triton.next_power_of_2(mtp_step + 1), + num_warps=1, + num_stages=1, + ) + expected = _reference_cumprod_probs(raw_probs, b_req_idx, mtp_step) + row = probs[0, : mtp_step + 1] + assert torch.allclose(row, expected[0, : mtp_step + 1], rtol=1e-5, atol=1e-5) + # Draft probabilities 0.995 and 0.005 clamp to 0.99 and 0.01. + assert row[0].item() == pytest.approx(1.0) + assert row[1].item() == pytest.approx(0.99, rel=1e-4) + assert row[2].item() == pytest.approx(0.99 * 0.01, rel=1e-4) + assert row[3].item() == pytest.approx(0.99 * 0.01 * 0.5, rel=1e-4) + + +def test_sample_select_count(): + mtp_step = 3 + req_num = 3 + probs, b_req_idx = _make_batch_probs( + req_num, + mtp_step, + rows=[ + [1.0, 0.95, 0.90, 0.10], + [1.0, 0.20, 0.80, 0.80], + [1.0, 0.99, 0.99, 0.99], + ], + ) + all_num = req_num * (mtp_step + 1) + for dynamic_batch_size in [3, 8, all_num]: + select = sample_dynamic_mtp_req_mask( + dynamic_batch_size=dynamic_batch_size, + b_req_idx=b_req_idx, + req_to_next_token_probs=probs.clone(), + mtp_step=mtp_step, + ) + assert select.dtype == torch.int32 + assert select.shape[0] == all_num + assert int(select.sum().item()) == dynamic_batch_size + assert torch.all((select == 0) | (select == 1)) + + +def test_sample_accepts_numpy_scalar_dynamic_batch_size(): + mtp_step = 3 + probs, b_req_idx = _make_batch_probs( + 3, + mtp_step, + rows=[ + [1.0, 0.95, 0.90, 0.10], + [1.0, 0.20, 0.80, 0.80], + [1.0, 0.99, 0.99, 0.99], + ], + ) + select = sample_dynamic_mtp_req_mask( + dynamic_batch_size=np.int64(8), + b_req_idx=b_req_idx, + req_to_next_token_probs=probs, + mtp_step=np.int64(mtp_step), + ) + assert int(select.sum().item()) == 8 + + +def test_sample_topk_by_cumprod_score(): + mtp_step = 3 + probs, b_req_idx = _make_batch_probs( + 3, + mtp_step, + rows=[ + [1.0, 0.95, 0.90, 0.10], + [1.0, 0.20, 0.80, 0.80], + [1.0, 0.99, 0.99, 0.99], + ], + ) + flat_probs = _flat_cumprod_probs(b_req_idx, probs, mtp_step) + for dynamic_batch_size in [1, 4, 8, 12]: + select = sample_dynamic_mtp_req_mask( + dynamic_batch_size=dynamic_batch_size, + b_req_idx=b_req_idx, + req_to_next_token_probs=probs.clone(), + mtp_step=mtp_step, + ) + _assert_topk_mask(select, flat_probs, dynamic_batch_size) + + +def test_sample_picks_highest_cumprod_rows(): + mtp_step = 1 + probs, b_req_idx = _make_batch_probs( + 2, + mtp_step, + rows=[ + [1.0, 0.9], + [1.0, 0.1], + ], + ) + flat_probs = _flat_cumprod_probs(b_req_idx, probs, mtp_step) + select = sample_dynamic_mtp_req_mask( + dynamic_batch_size=2, + b_req_idx=b_req_idx, + req_to_next_token_probs=probs.clone(), + mtp_step=mtp_step, + ) + _assert_topk_mask(select, flat_probs, 2) + # top-2 scores are both 0.99 at mtp_index==0 (req0 and req1 main rows) + assert select[0].item() == 1 + assert select[2].item() == 1 + + +def test_sample_single_request(): + mtp_step = 2 + probs, b_req_idx = _make_batch_probs(1, mtp_step, rows=[[1.0, 0.5, 0.25]]) + flat_probs = _flat_cumprod_probs(b_req_idx, probs, mtp_step) + select = sample_dynamic_mtp_req_mask( + dynamic_batch_size=2, + b_req_idx=b_req_idx, + req_to_next_token_probs=probs.clone(), + mtp_step=mtp_step, + ) + _assert_topk_mask(select, flat_probs, 2) + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py b/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py new file mode 100644 index 0000000000..b28ed9417b --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py @@ -0,0 +1,82 @@ +import pytest +import torch + +if not torch.cuda.is_available(): + pytest.skip("requires CUDA", allow_module_level=True) + +from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_mtp_fa3_decode_params + + +def _reference_dynamic_mtp_fa3_decode_params(b_req_idx, b_seq_len, b_mark_shared_group, hold_req_id): + batch_size = b_req_idx.shape[0] + b_req_idx_cpu = b_req_idx.cpu() + b_seq_len_cpu = b_seq_len.cpu() + b_mark_cpu = b_mark_shared_group.cpu() + valid_idx = torch.nonzero(b_mark_cpu > 0, as_tuple=False).flatten() + valid_size = valid_idx.numel() + + b_q_seq_len = torch.zeros((batch_size,), dtype=torch.int32) + b_kv_seq_len = torch.zeros((batch_size,), dtype=torch.int32) + b_att_req_idx = torch.full((batch_size,), hold_req_id, dtype=torch.int32) + b_att_seq_len = torch.zeros((batch_size,), dtype=torch.int32) + + b_q_seq_len[:valid_size] = b_mark_cpu[valid_idx] + b_kv_seq_len[:valid_size] = b_seq_len_cpu[valid_idx] + b_att_req_idx[:valid_size] = b_req_idx_cpu[valid_idx] + b_att_seq_len[:valid_size] = b_seq_len_cpu[valid_idx] + return b_q_seq_len, b_kv_seq_len, b_att_req_idx, b_att_seq_len + + +@pytest.mark.parametrize("batch_size", [1, 7, 256, 257, 777, 1025, 1537]) +def test_build_dynamic_mtp_fa3_decode_params(batch_size): + hold_req_id = -1 + b_req_idx = torch.arange(batch_size, dtype=torch.int32, device="cuda") + 100 + b_seq_len = (torch.arange(batch_size, dtype=torch.int32, device="cuda") % 97) + 1 + b_mark_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device="cuda") + + mark_positions = [0, 3, 17, 255, 256, 511, batch_size - 1] + for pos in mark_positions: + if 0 <= pos < batch_size: + b_mark_shared_group[pos] = pos % 5 + 1 + + actual = build_dynamic_mtp_fa3_decode_params( + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + att_batch_size=batch_size, + hold_req_id=hold_req_id, + ) + expected = _reference_dynamic_mtp_fa3_decode_params( + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + hold_req_id=hold_req_id, + ) + + for actual_tensor, expected_tensor in zip(actual, expected, strict=True): + assert torch.equal(actual_tensor.cpu(), expected_tensor) + + +def test_build_dynamic_mtp_fa3_decode_params_all_padding(): + batch_size = 513 + hold_req_id = -1 + b_req_idx = torch.arange(batch_size, dtype=torch.int32, device="cuda") + b_seq_len = torch.arange(batch_size, dtype=torch.int32, device="cuda") + 1 + b_mark_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device="cuda") + + actual = build_dynamic_mtp_fa3_decode_params( + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + att_batch_size=batch_size, + hold_req_id=hold_req_id, + ) + expected = _reference_dynamic_mtp_fa3_decode_params( + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + b_mark_shared_group=b_mark_shared_group, + hold_req_id=hold_req_id, + ) + + for actual_tensor, expected_tensor in zip(actual, expected, strict=True): + assert torch.equal(actual_tensor.cpu(), expected_tensor) diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py new file mode 100644 index 0000000000..42b2bfab32 --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -0,0 +1,166 @@ +import json +import pytest +import torch + +if not torch.cuda.is_available(): + pytest.skip("requires CUDA", allow_module_level=True) + +from lightllm.common.basemodel.batch_objs import ModelInput +from lightllm.common.basemodel.triton_kernel import mtp_utils + + +def test_trim_dynamic_mtp_model_input(monkeypatch): + monkeypatch.setenv( + "LIGHTLLM_START_ARGS", + json.dumps( + { + "diverse_mode": False, + "llm_kv_type": "fp16", + "mtp_dynamic_verify": True, + "mtp_step": 3, + "llm_decode_att_backend": "triton", + } + ), + ) + monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "4") + mtp_utils.get_env_start_args.cache_clear() + mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() + + model_input = ModelInput( + batch_size=12, + total_token_num=54, + max_q_seq_len=1, + max_kv_seq_len=6, + input_ids=torch.arange(12, dtype=torch.int64, device="cuda") + 1000, + b_req_idx=torch.tensor([0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2], dtype=torch.int32, device="cuda"), + b_mtp_index=torch.tensor([0, 1, 2, 3, 0, 1, 2, 3, 0, 1, 2, 3], dtype=torch.int32, device="cuda"), + b_seq_len=torch.tensor([3, 4, 5, 6, 3, 4, 5, 6, 3, 4, 5, 6], dtype=torch.int32, device="cuda"), + b_position_delta=torch.arange(12, dtype=torch.int32, device="cuda") + 200, + b_shared_seq_len=torch.tensor([0, 0, 0, 0, 7, 7, 7, 7, 9, 9, 9, 9], dtype=torch.int32, device="cuda"), + b_mark_shared_group=torch.tensor([0, 0, 0, 4, 0, 0, 0, 4, 0, 0, 0, 4], dtype=torch.int32, device="cuda"), + mem_indexes=torch.arange(12, dtype=torch.int32, device="cuda") + 100, + mem_indexes_cpu=torch.arange(12, dtype=torch.int32, device="cpu") + 100, + is_prefill=False, + multimodal_params=[{"row": i, "images": [], "audios": []} for i in range(12)], + mtp_draft_input_hiddens=(torch.arange(12 * 5, dtype=torch.float32, device="cuda").reshape(12, 5) + 0.5), + ) + req_to_next_token_probs = torch.tensor( + [ + [1.0, 0.95, 0.90, 0.10, 0.0, 0.0], + [1.0, 0.20, 0.80, 0.80, 0.0, 0.0], + [1.0, 0.99, 0.99, 0.99, 0.0, 0.0], + ], + dtype=torch.float32, + device="cuda", + ) + + trimmed_input, selected_mask = mtp_utils.prepare_dynamic_mtp_model_input( + model_input=model_input, + req_num=3, + dynamic_batch_size=8, + req_to_next_token_ids=torch.empty((0,), dtype=torch.int64, device="cuda"), + req_to_next_token_probs=req_to_next_token_probs, + ) + torch.cuda.synchronize() + + expected_selected_mask = torch.tensor([1, 1, 1, 0, 1, 0, 0, 0, 1, 1, 1, 1], dtype=torch.int32) + expected_selected_rows = torch.where(expected_selected_mask == 1)[0] + + assert torch.equal(selected_mask.cpu(), expected_selected_mask) + assert trimmed_input.batch_size == 8 + assert trimmed_input.max_q_seq_len == 1 + assert trimmed_input.multimodal_params == [{"images": [], "audios": []}] * 8 + + assert torch.equal( + trimmed_input.input_ids.cpu(), torch.arange(12, dtype=torch.int64)[expected_selected_rows] + 1000 + ) + assert torch.equal(trimmed_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 1, 2, 2, 2, 2], dtype=torch.int32)) + assert torch.equal(trimmed_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 0, 0, 1, 2, 3], dtype=torch.int32)) + assert torch.equal(trimmed_input.b_seq_len.cpu(), torch.tensor([3, 4, 5, 3, 3, 4, 5, 6], dtype=torch.int32)) + assert torch.equal( + trimmed_input.b_position_delta.cpu(), + torch.tensor([200, 201, 202, 204, 208, 209, 210, 211], dtype=torch.int32), + ) + assert torch.equal(trimmed_input.b_shared_seq_len.cpu(), torch.tensor([0, 0, 0, 7, 9, 9, 9, 9], dtype=torch.int32)) + assert torch.equal( + trimmed_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 1, 0, 0, 0, 4], dtype=torch.int32) + ) + assert torch.equal( + trimmed_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 104, 108, 109, 110, 111], dtype=torch.int32) + ) + # The hot path intentionally keeps the CPU copy unfiltered to avoid a GPU-to-CPU + # synchronization. The router frees rejected indexes after its async mask copy. + assert torch.equal(trimmed_input.mem_indexes_cpu, torch.arange(12, dtype=torch.int32) + 100) + + expected_hiddens = (torch.arange(12 * 5, dtype=torch.float32).reshape(12, 5) + 0.5)[expected_selected_rows] + assert torch.equal(trimmed_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) + + +def test_trim_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size(monkeypatch): + monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "3") + mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() + + model_input = ModelInput( + batch_size=5, + total_token_num=36, + max_q_seq_len=1, + max_kv_seq_len=6, + input_ids=None, + b_req_idx=torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32, device="cuda"), + b_mtp_index=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device="cuda"), + b_seq_len=torch.tensor([3, 4, 5, 6, 7], dtype=torch.int32, device="cuda"), + b_shared_seq_len=None, + b_mark_shared_group=torch.tensor([0, 0, 0, 0, 5], dtype=torch.int32, device="cuda"), + mem_indexes=torch.arange(5, dtype=torch.int32, device="cuda"), + mem_indexes_cpu=torch.arange(5, dtype=torch.int32, device="cpu"), + is_prefill=False, + multimodal_params=[{"images": [], "audios": []} for _ in range(5)], + ) + selected_mask = torch.ones((5,), dtype=torch.int32, device="cuda") + + trimmed_input = mtp_utils._trim_decode_model_input_inplace( + model_input=model_input, + selected_mask_gpu=selected_mask, + dynamic_batch_size=5, + ) + torch.cuda.synchronize() + + assert torch.equal(trimmed_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32)) + assert torch.equal(trimmed_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32)) + assert torch.equal(trimmed_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 0, 2], dtype=torch.int32)) + + +def test_mtp_verify_scatter_and_start_locations(): + req_to_next_token_ids = torch.tensor( + [[1, 2, -2, -1, -1], [1, 2, 0, -1, -1], [1, 3, 4, 4, 5]], + dtype=torch.int64, + device="cuda", + ) + b_req_idx = torch.tensor([0, 0, 2, 2, 2], dtype=torch.int32, device="cuda") + b_mtp_index = torch.tensor([0, 1, 0, 1, 2], dtype=torch.int32, device="cuda") + b_req_mtp_start_loc = mtp_utils.gen_b_req_mtp_start_loc(b_mtp_index, num_reqs=2) + new_next_token_ids = torch.tensor([1, 4, 2, 4, 13], dtype=torch.int64, device="cuda") + all_next_token_ids = torch.tensor( + [[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], [13, 14, 15]], + dtype=torch.int64, + device="cuda", + ) + + mtp_accept_len, accepted_index = mtp_utils.mtp_verify( + req_to_next_token_ids, b_req_mtp_start_loc, new_next_token_ids, b_req_idx + ) + mtp_utils.mtp_scatter_next_token_ids( + req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, mtp_accept_len + ) + torch.cuda.synchronize() + + assert torch.equal(b_req_mtp_start_loc.cpu(), torch.tensor([0, 2], dtype=torch.int64)) + assert torch.equal(mtp_accept_len.cpu(), torch.tensor([1, 1], dtype=torch.int32)) + assert torch.equal(accepted_index.cpu(), torch.tensor([1, 0, 1, 0, 0], dtype=torch.int32)) + assert torch.equal( + req_to_next_token_ids.cpu(), + torch.tensor( + [[1, 2, 3, -1, -1], [1, 2, 0, -1, -1], [7, 8, 9, 4, 5]], + dtype=torch.int64, + ), + ) diff --git a/unit_tests/common/speculative/test_config.py b/unit_tests/common/speculative/test_config.py new file mode 100644 index 0000000000..cbcd4d6394 --- /dev/null +++ b/unit_tests/common/speculative/test_config.py @@ -0,0 +1,40 @@ +from types import SimpleNamespace + +import pytest + +from lightllm.common.speculative.config import SpeculativeConfig + + +@pytest.mark.parametrize( + ("mode", "step", "requested_dynamic", "expected_dynamic", "draft_model_count"), + [ + (None, 0, False, False, 0), + ("vanilla_with_att", 3, False, False, 3), + ("eagle3", 3, True, True, 1), + ("dspark", 0, False, True, 1), + ("dflash", 0, True, False, 1), + ], +) +def test_speculative_config_normalizes_modes( + mode, step, requested_dynamic, expected_dynamic, draft_model_count +): + config = SpeculativeConfig.from_args( + SimpleNamespace(mtp_mode=mode, mtp_step=step, mtp_dynamic_verify=requested_dynamic) + ) + + config.validate() + assert config.dynamic_verify is expected_dynamic + assert config.draft_model_count == draft_model_count + assert config.enabled is (mode is not None) + + +def test_dynamic_eagle3_uses_unit_graph_granularity(): + config = SpeculativeConfig(mode="eagle3", step=7, dynamic_verify=True) + + assert config.get_decode_graph_mtp_step(model_config={}, is_draft_model=False) == 0 + assert config.get_decode_graph_warmup_mtp_step(model_config={}, is_draft_model=False) == 3 + + +def test_invalid_mode_is_rejected(): + with pytest.raises(AssertionError, match="unsupported speculative mode"): + SpeculativeConfig(mode="unknown", step=1).validate() diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py new file mode 100644 index 0000000000..e8b3c19719 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py @@ -0,0 +1,50 @@ +import pytest +import torch + +if not torch.cuda.is_available(): + pytest.skip("requires CUDA", allow_module_level=True) + +from lightllm.server.router.model_infer.mode_backend.generic_post_process import _trim_post_sample_tensors + + +def test_trim_post_sample_tensors(): + selected = torch.tensor([1, 1, 1, 0, 1, 0, 0, 0, 1, 1, 1, 1], dtype=torch.int32, device="cuda") + selected_rows = torch.where(selected.cpu() == 1)[0] + dynamic_batch_size = int(selected.sum().item()) + + b_req_idx = torch.arange(12, dtype=torch.int32, device="cuda") + 10 + b_temperatures = torch.arange(12, dtype=torch.float32, device="cuda") + 0.5 + b_top_ps = torch.arange(12, dtype=torch.float32, device="cuda") / 100 + 0.8 + b_top_ks = torch.arange(12, dtype=torch.int32, device="cuda") + 100 + b_length_penalty_param = torch.arange(12, dtype=torch.int32, device="cuda") + 200 + b_mask_eos_reqs = torch.tensor( + [True, False, True, False, True, False, True, False, True, False, True, False], + dtype=torch.bool, + device="cuda", + ) + + ( + out_b_req_idx, + out_b_temperatures, + out_b_top_ps, + out_b_top_ks, + out_b_length_penalty_param, + out_b_mask_eos_reqs, + ) = _trim_post_sample_tensors( + dynamic_batch_size=dynamic_batch_size, + selected_run_reqs=selected, + b_req_idx=b_req_idx, + b_temperatures=b_temperatures, + b_top_ps=b_top_ps, + b_top_ks=b_top_ks, + b_length_penalty_param=b_length_penalty_param, + b_mask_eos_reqs=b_mask_eos_reqs, + ) + torch.cuda.synchronize() + + assert torch.equal(out_b_req_idx.cpu(), b_req_idx.cpu()[selected_rows]) + assert torch.equal(out_b_temperatures.cpu(), b_temperatures.cpu()[selected_rows]) + assert torch.equal(out_b_top_ps.cpu(), b_top_ps.cpu()[selected_rows]) + assert torch.equal(out_b_top_ks.cpu(), b_top_ks.cpu()[selected_rows]) + assert torch.equal(out_b_length_penalty_param.cpu(), b_length_penalty_param.cpu()[selected_rows]) + assert torch.equal(out_b_mask_eos_reqs.cpu(), b_mask_eos_reqs.cpu()[selected_rows]) diff --git a/unit_tests/server/router/model_infer/speculative/test_planner.py b/unit_tests/server/router/model_infer/speculative/test_planner.py new file mode 100644 index 0000000000..b5987e1396 --- /dev/null +++ b/unit_tests/server/router/model_infer/speculative/test_planner.py @@ -0,0 +1,66 @@ +import random + +import numpy as np + +from lightllm.server.router.model_infer.speculative.planner import ( + DSparkDynamicMTPPlanner, + DynamicMTPPlanner, + FixedMTPPlanner, +) + + +def test_fixed_planner_returns_static_plan(): + plan = FixedMTPPlanner(mtp_step=3).plan(req_num=4, original_batch_size=16) + + assert not plan.is_dynamic + assert plan.dynamic_batch_size is None + assert plan.draft_step == plan.pre_draft_step == 3 + assert plan.selection_mode == "none" + assert not plan.skip_verify_sync + + +def test_dynamic_planner_stays_full_width_until_costs_are_profiled(): + plan = DynamicMTPPlanner(mtp_step=3, use_random_mode=False).plan( + req_num=2, original_batch_size=8 + ) + + assert plan.dynamic_batch_size == 8 + assert plan.draft_step == plan.pre_draft_step == 3 + assert plan.selection_mode == "observe" + + +def test_planner_does_not_reset_process_global_random_state(): + random.seed(2027) + expected = random.random() + random.seed(2027) + + DynamicMTPPlanner(mtp_step=3) + + assert random.random() == expected + + +def test_dspark_applies_confidence_capacity_after_two_step_delay(): + planner = DSparkDynamicMTPPlanner(mtp_step=3) + planner.update_infer_cost(batch_size=2, infer_cost_ms=1.0, is_draft_model=False) + planner.update_infer_cost(batch_size=4, infer_cost_ms=1.1, is_draft_model=False) + planner.update_infer_cost(batch_size=8, infer_cost_ms=10.0, is_draft_model=False) + schedule_probs = np.asarray([[1.0, 0.9, 0.9, 0.9]] * 2, dtype=np.float64) + + planner.update_predicted_schedule_probs(schedule_probs=schedule_probs, req_num=2) + first_plan = planner.plan(req_num=2, original_batch_size=8) + planner.update_predicted_schedule_probs(schedule_probs=schedule_probs, req_num=2) + second_plan = planner.plan(req_num=2, original_batch_size=8) + + assert first_plan.dynamic_batch_size == 2 + assert second_plan.dynamic_batch_size == 4 + assert second_plan.draft_step == 3 + + +def test_topk_prefix_sums_only_computes_requested_counts(): + result = DSparkDynamicMTPPlanner._topk_prefix_sums( + values=np.asarray([0.1, 0.9, 0.4, 0.7]), + counts=[0, 2, 4], + ) + + assert set(result) == {0, 2, 4} + np.testing.assert_allclose([result[0], result[2], result[4]], [0.0, 1.6, 2.1]) diff --git a/unit_tests/server/test_api_start_spec_config.py b/unit_tests/server/test_api_start_spec_config.py new file mode 100644 index 0000000000..8982969703 --- /dev/null +++ b/unit_tests/server/test_api_start_spec_config.py @@ -0,0 +1,42 @@ +from types import SimpleNamespace + +from lightllm.common.speculative.config import SpeculativeConfig +from lightllm.server import api_start + + +def test_block_mode_derives_mtp_step_from_checkpoint(monkeypatch): + draft_config = { + "architectures": ["Qwen3DSparkModel"], + "block_size": 4, + "target_layer_ids": [8, 16, 24], + "mask_token_id": 151665, + "markov_rank": 0, + "enable_confidence_head": True, + "confidence_head_with_markov": False, + } + monkeypatch.setattr( + api_start.PretrainedConfig, + "get_config_dict", + lambda _model_dir: (draft_config, {}), + ) + args = SimpleNamespace(mtp_draft_model_dir=["/models/dspark"], mtp_step=0) + config = SpeculativeConfig(mode="dspark", step=0, dynamic_verify=True) + + normalized = api_start.normalize_block_mtp_step_from_first_draft_config(args, config) + + assert args.mtp_step == 4 + assert normalized.step == 4 + normalized.validate() + + +def test_non_block_mode_keeps_explicit_step(monkeypatch): + monkeypatch.setattr( + api_start.PretrainedConfig, + "get_config_dict", + lambda _model_dir: (_ for _ in ()).throw(AssertionError("must not load config")), + ) + args = SimpleNamespace(mtp_draft_model_dir=["/models/eagle3"], mtp_step=3) + config = SpeculativeConfig(mode="eagle3", step=3, dynamic_verify=True) + + assert api_start.normalize_block_mtp_step_from_first_draft_config(args, config) is config + assert args.mtp_step == 3 From 0b00906cdbedaac438d48309845f0cf9d695848b Mon Sep 17 00:00:00 2001 From: Zhao Xintong Date: Tue, 4 Aug 2026 07:22:22 +0000 Subject: [PATCH 002/103] add qwen3.5 dflash support --- lightllm/common/basemodel/basemodel.py | 4 - .../triton_kernel/linear_att_copy.py | 14 +- .../operator/linear_att.py | 38 ++ .../qwen3next_mem_manager.py | 165 +++++++- lightllm/common/speculative/__init__.py | 10 +- lightllm/common/speculative/config.py | 98 ++++- lightllm/models/__init__.py | 1 + lightllm/models/qwen3_5_dflash/__init__.py | 3 + .../qwen3_5_dflash/layer_weights/__init__.py | 5 + .../pre_and_post_layer_weight.py | 47 +++ lightllm/models/qwen3_5_dflash/model.py | 135 +++++++ .../layer_infer/transformer_layer_infer.py | 36 ++ lightllm/models/qwen3_dflash/model.py | 1 + lightllm/server/api_start.py | 20 +- lightllm/server/pd_io_struct.py | 2 +- .../server/router/model_infer/infer_batch.py | 1 + .../model_infer/mode_backend/base_backend.py | 67 +++- .../pd/decode_node_impl/decode_impl.py | 11 +- .../pd/prefill_node_impl/prefill_impl.py | 24 +- .../speculative/proposers/dflash.py | 51 ++- .../speculative/proposers/dspark.py | 22 +- .../router/model_infer/speculative/runtime.py | 3 +- lightllm/utils/envs_utils.py | 6 +- lightllm/utils/kv_cache_utils.py | 41 +- test/speculative/test_qwen35_dflash_state.py | 378 ++++++++++++++++++ 25 files changed, 1108 insertions(+), 75 deletions(-) create mode 100644 lightllm/models/qwen3_5_dflash/__init__.py create mode 100644 lightllm/models/qwen3_5_dflash/layer_weights/__init__.py create mode 100644 lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py create mode 100644 lightllm/models/qwen3_5_dflash/model.py create mode 100644 test/speculative/test_qwen35_dflash_state.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index ab97b81469..3199082c66 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -596,10 +596,6 @@ def _create_unpad_prefill_model_output( if new_model_output.prompt_logics is not None: new_model_output.prompt_logics = new_model_output.prompt_logics[0:origin_handle_token_num] - # 特殊模型,特殊模式的特殊变量的特殊 unpad - if new_model_output.mtp_main_output_hiddens is not None: - _hidden_states = new_model_output.mtp_main_output_hiddens - new_model_output.mtp_main_output_hiddens = _hidden_states[0:origin_handle_token_num] if self.spec_adapter is not None: self.spec_adapter.unpad_hidden(token_num=origin_handle_token_num, microbatch_index=microbatch_index) diff --git a/lightllm/common/basemodel/triton_kernel/linear_att_copy.py b/lightllm/common/basemodel/triton_kernel/linear_att_copy.py index eb17507fb0..0f2841ca2f 100644 --- a/lightllm/common/basemodel/triton_kernel/linear_att_copy.py +++ b/lightllm/common/basemodel/triton_kernel/linear_att_copy.py @@ -10,6 +10,7 @@ def _copy_linear_att_state_to_kv_buffer( cpu_kv_conv_ptr, # uint8 view: [buffer_num, linear_layer_num, conv_dim * cpu_conv_row_bytes] cpu_kv_ssm_ptr, # uint8 view: [buffer_num, linear_layer_num, ssm_bytes] b_req_idx, # [batch_size,] + req_to_mtp_state_index, # [max_request_num + 1,] big_page_buffer_ids, # [batch_size,] gpu_conv_stride_l, gpu_conv_stride_s, @@ -24,7 +25,7 @@ def _copy_linear_att_state_to_kv_buffer( cpu_kv_ssm_stride_s, cpu_kv_ssm_stride_l, cpu_kv_ssm_stride_d, - mtp_step, + mtp_step: tl.constexpr, gpu_conv_dim, # number of conv rows gpu_conv_tail_dim_bytes, # bytes copied per conv row; equals the CPU/cache row width gpu_ssm_tail_dim, @@ -50,7 +51,10 @@ def _copy_linear_att_state_to_kv_buffer( return cur_req_idx = tl.load(b_req_idx + cur_batch).to(tl.int64) - cur_state_req_idx = (cur_req_idx * (mtp_step + 1)).to(tl.int64) + state_offset = 0 + if mtp_step > 0: + state_offset = tl.load(req_to_mtp_state_index + cur_req_idx).to(tl.int64) + cur_state_req_idx = (cur_req_idx * (mtp_step + 1) + state_offset).to(tl.int64) gpu_conv_base = gpu_conv_ptr + cur_layer * gpu_conv_stride_l + cur_req_idx * gpu_conv_stride_s cpu_conv_base = cpu_kv_conv_ptr + big_page_buffer_idx * cpu_kv_conv_stride_s + cur_layer * cpu_kv_conv_stride_l @@ -80,6 +84,7 @@ def _copy_linear_att_state_to_kv_buffer( def copy_linear_att_state_to_kv_buffer( b_req_idx: torch.Tensor, + req_to_mtp_state_index: torch.Tensor, big_page_buffer_ids: torch.Tensor, gpu_conv_state: torch.Tensor, # [linear_layer_num, req_num, conv_dim, kernel_size] gpu_ssm_state: torch.Tensor, # [linear_layer_num, req_num * (mtp_step + 1), ...] @@ -89,6 +94,10 @@ def copy_linear_att_state_to_kv_buffer( ): # gpu_conv_state 的后两维可能是不连续的。 assert len(b_req_idx) == big_page_buffer_ids.shape[0] + if req_to_mtp_state_index is None: + assert mtp_step == 0 + # The constexpr branch below does not dereference this placeholder. + req_to_mtp_state_index = b_req_idx BLOCK = 4096 assert gpu_conv_state.dim() == 4, "gpu_conv_state must be [layer, s, conv_dim, widened_width]" @@ -129,6 +138,7 @@ def copy_linear_att_state_to_kv_buffer( cpu_kv_conv_ptr=cpu_kv_conv_state, cpu_kv_ssm_ptr=cpu_kv_ssm_state, b_req_idx=b_req_idx, + req_to_mtp_state_index=req_to_mtp_state_index, big_page_buffer_ids=big_page_buffer_ids, gpu_conv_stride_l=gpu_conv_state.stride(0), gpu_conv_stride_s=gpu_conv_state.stride(1), diff --git a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py index 147b43f697..6d5152e4a6 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py +++ b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py @@ -89,6 +89,24 @@ def load_cpu_cache_to_gpu( big_page_token_num=args.cpu_cache_token_page_size, linear_config=self.linear_config, ) + if mem_manager.has_separate_dflash_draft_kv: + from lightllm.common.basemodel.triton_kernel.kv_cache_offload import ( + load_cpu_kv_to_gpu, + ) + + draft_mem = mem_manager._pd_dflash_draft_mem_manager + draft_mem_indexes = mem_indexes.masked_fill(mem_indexes < 0, draft_mem.HOLD_TOKEN_MEMINDEX) + load_cpu_kv_to_gpu( + gpu_mem_indexes=draft_mem_indexes, + gpu_kv_cache=draft_mem.kv_buffer, + gpu_kv_cache_scale=None, + cpu_kv_cache=mem_manager.get_dflash_draft_cpu_cache(cpu_cache_client.cpu_kv_cache_tensor), + cpu_kv_cache_scale=None, + page_indexes=page_indexes, + tp_index=get_current_rank_in_dp(), + tp_world_size=get_dp_world_size(), + grid_num=16, + ) from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -187,6 +205,26 @@ def offload_gpu_kv_to_cpu_cache( big_page_token_num=args.cpu_cache_token_page_size, linear_config=self.linear_config, ) + if mem_manager.has_separate_dflash_draft_kv: + from lightllm.common.basemodel.triton_kernel.kv_cache_offload import ( + offload_gpu_kv_to_cpu, + ) + + draft_mem = mem_manager._pd_dflash_draft_mem_manager + draft_mem_indexes = mem_indexes.masked_fill(mem_indexes < 0, draft_mem.HOLD_TOKEN_MEMINDEX) + offload_gpu_kv_to_cpu( + token_indexes=draft_mem_indexes, + gpu_kv_cache=draft_mem.kv_buffer, + gpu_kv_cache_scale=None, + cpu_kv_cache=mem_manager.get_dflash_draft_cpu_cache(cpu_cache_client.cpu_kv_cache_tensor), + cpu_kv_cache_scale=None, + page_indexes=page_indexes, + page_readies=page_readies, + tp_index=get_current_rank_in_dp(), + tp_world_size=get_dp_world_size(), + grid_num=16, + ) + return def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): diff --git a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py index c8f69dac84..d37a226ea9 100644 --- a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py @@ -2,6 +2,7 @@ import triton from lightllm.utils.log_utils import init_logger from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager +from lightllm.common.kv_trans_kernel.nixl_kv_trans import page_io from lightllm.utils.envs_utils import get_env_start_args from lightllm.common.linear_att_cache_manager import LinearAttCacheConfig, LinearAttCacheManager from .operator import LinearAttMemOperator @@ -25,9 +26,25 @@ def __init__( mem_fraction=0.9, ): self.linear_config = linear_config + self._pd_dflash_draft_mem_manager = None + self._pd_dflash_global_kv_heads = None super().__init__(size, dtype, num_kv_heads, head_dim, full_att_layer_num, always_copy, mem_fraction) + @property + def has_separate_dflash_draft_kv(self) -> bool: + return self._pd_dflash_draft_mem_manager is not None + + def register_dflash_draft_mem_manager(self, draft_mem_manager: MemoryManager, global_kv_heads: int): + """Attach the independent Qwen3.5 DFlash KV cache for PD transfer.""" + assert draft_mem_manager.size == self.size + assert draft_mem_manager.dtype == self.dtype + global_kv_heads = int(global_kv_heads) + assert global_kv_heads > 0 + self._pd_dflash_draft_mem_manager = draft_mem_manager + self._pd_dflash_global_kv_heads = global_kv_heads + return + def get_att_input_params(self, layer_index: int) -> Tuple[Any, Any]: if layer_index >= self.linear_config.all_layer_num: # MTP draft full-attn layers are packed after the main model layers. @@ -84,10 +101,146 @@ def write_to_shm(self, req_manager): self.linear_att_big_page_buffers = big_page_buffers def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: - kv_move_buffer = super().alloc_paged_kv_move_buffer(page_num, page_size) + if not self.has_separate_dflash_draft_kv: + kv_move_buffer = super().alloc_paged_kv_move_buffer(page_num, page_size) + else: + # PD pages are opaque contiguous bytes to the transporter. Allocate + # one registered page large enough for either target or draft KV; + # the page-kind hooks below expose the corresponding shaped view. + target_elements = ( + self.layer_num * 2 * self.linear_config.full_att_all_num_kv_heads * self.head_dim + ) + draft_mem = self._pd_dflash_draft_mem_manager + draft_elements = ( + draft_mem.layer_num * 2 * self._pd_dflash_global_kv_heads * draft_mem.head_dim + ) + linear_state_nbytes = Qwen3NextLinearAttPageHelper(self).state_nbytes + linear_state_elements_per_token = ( + linear_state_nbytes + page_size * self.dtype.itemsize - 1 + ) // (page_size * self.dtype.itemsize) + elements_per_token = max( + target_elements, + draft_elements, + linear_state_elements_per_token, + ) + self.kv_move_buffer = torch.empty( + (page_num, page_size, 1, 1, elements_per_token), + dtype=self.dtype, + device="cuda", + ) + self._buffer_mem_indexes_tensors = [ + torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) + for _ in range(page_num) + ] + kv_move_buffer = self.kv_move_buffer Qwen3NextLinearAttPageHelper(self).assert_page_size() return kv_move_buffer + def _get_pd_kv_page(self, page_index: int, page_kind: str) -> torch.Tensor: + page = self.kv_move_buffer[page_index] + if not self.has_separate_dflash_draft_kv: + assert page_kind == "kv" + return page + + flat_page = page.reshape(-1) + if page_kind == "kv": + layer_num = self.layer_num + global_kv_heads = self.linear_config.full_att_all_num_kv_heads + head_dim = self.head_dim + elif page_kind == "draft_kv": + draft_mem = self._pd_dflash_draft_mem_manager + layer_num = draft_mem.layer_num + global_kv_heads = self._pd_dflash_global_kv_heads + head_dim = draft_mem.head_dim + else: + raise ValueError(f"unknown KV page kind {page_kind}") + + element_num = layer_num * 2 * global_kv_heads * head_dim + total_element_num = page.shape[0] * element_num + assert total_element_num <= flat_page.numel() + return flat_page[:total_element_num].view(page.shape[0], layer_num, 2 * global_kv_heads, head_dim) + + def _get_pd_mem_manager(self, page_kind: str) -> MemoryManager: + if page_kind == "kv": + return self + if page_kind == "draft_kv" and self.has_separate_dflash_draft_kv: + return self._pd_dflash_draft_mem_manager + raise ValueError(f"unknown KV page kind {page_kind}") + + def get_dflash_draft_cpu_cache(self, cpu_cache_tensor: torch.Tensor) -> torch.Tensor: + """Return the draft-KV portion appended to a mixed-model CPU cache page.""" + + assert self.has_separate_dflash_draft_kv + assert cpu_cache_tensor.dtype == torch.uint8 + draft_mem = self._pd_dflash_draft_mem_manager + page_num = cpu_cache_tensor.shape[0] + page_size = get_env_start_args().cpu_cache_token_page_size + draft_shape = ( + page_num, + draft_mem.layer_num, + page_size, + 2 * self._pd_dflash_global_kv_heads, + draft_mem.head_dim, + ) + draft_nbytes = 1 + for dim in draft_shape[1:]: + draft_nbytes *= dim + draft_nbytes *= self.dtype.itemsize + + # Existing Qwen3Next full-attention/linear-state data occupies the + # aligned prefix. Qwen3.5 DFlash KV is stored in the remaining bytes so + # disk offload and shared-memory page lifecycle stay unchanged. + draft_offset = self.linear_config.get_cpu_cache_big_page_bytes() + flat_cache = cpu_cache_tensor.reshape(page_num, -1) + assert draft_offset + draft_nbytes <= flat_cache.shape[1] + draft_bytes = flat_cache[:, draft_offset : draft_offset + draft_nbytes] + return draft_bytes.view(dtype=self.dtype).view(draft_shape) + + def _write_kv_mem_to_page( + self, mem_indexes, page_index: int, dp_index: int, mem_managers, dp_world_size: int, page_kind: str + ): + page = self._get_pd_kv_page(page_index, page_kind) + pin_mem_indexes = self._buffer_mem_indexes_tensors[page_index][0 : len(mem_indexes)] + pin_mem_indexes.numpy()[:] = mem_indexes + mem_indexes_gpu = pin_mem_indexes.cuda(non_blocking=True) + dp_mems = mem_managers[(dp_index * dp_world_size) : ((dp_index + 1) * dp_world_size)] + assert len(dp_mems) == dp_world_size + source_mems = [mem._get_pd_mem_manager(page_kind) for mem in dp_mems] + repeat_count = dp_world_size * source_mems[0].kv_buffer.shape[2] // page.shape[2] + assert repeat_count > 0 + for tp_index, mem in enumerate(source_mems): + if tp_index % repeat_count == 0: + page_io( + mem_indexes=mem_indexes_gpu, + page_tensor=page, + kv_buffer=mem.kv_buffer, + tp_index=tp_index, + tp_world_size=dp_world_size, + mode="write", + ) + return + + def _read_kv_page_to_mem( + self, mem_indexes, page_index: int, dp_index: int, mem_managers, dp_world_size: int, page_kind: str + ): + page = self._get_pd_kv_page(page_index, page_kind) + pin_mem_indexes = self._buffer_mem_indexes_tensors[page_index][0 : len(mem_indexes)] + pin_mem_indexes.numpy()[:] = mem_indexes + mem_indexes_gpu = pin_mem_indexes.cuda(non_blocking=True) + dp_mems = mem_managers[(dp_index * dp_world_size) : ((dp_index + 1) * dp_world_size)] + assert len(dp_mems) == dp_world_size + target_mems = [mem._get_pd_mem_manager(page_kind) for mem in dp_mems] + for tp_index, mem in enumerate(target_mems): + page_io( + mem_indexes=mem_indexes_gpu, + page_tensor=page, + kv_buffer=mem.kv_buffer, + tp_index=tp_index, + tp_world_size=dp_world_size, + mode="read", + ) + return + def write_mem_to_page_kv_move_buffer( self, mem_indexes, @@ -98,15 +251,14 @@ def write_mem_to_page_kv_move_buffer( page_kind: str = "kv", req_idx: int = None, ): - if page_kind == "kv": - return super().write_mem_to_page_kv_move_buffer( + if page_kind in ("kv", "draft_kv"): + return self._write_kv_mem_to_page( mem_indexes=mem_indexes, page_index=page_index, dp_index=dp_index, mem_managers=mem_managers, dp_world_size=dp_world_size, page_kind=page_kind, - req_idx=req_idx, ) assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" assert req_idx is not None @@ -125,15 +277,14 @@ def read_page_kv_move_buffer_to_mem( page_kind: str = "kv", req_idx: int = None, ): - if page_kind == "kv": - return super().read_page_kv_move_buffer_to_mem( + if page_kind in ("kv", "draft_kv"): + return self._read_kv_page_to_mem( mem_indexes=mem_indexes, page_index=page_index, dp_index=dp_index, mem_managers=mem_managers, dp_world_size=dp_world_size, page_kind=page_kind, - req_idx=req_idx, ) assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" assert req_idx is not None diff --git a/lightllm/common/speculative/__init__.py b/lightllm/common/speculative/__init__.py index 6507ee69ab..1388ca7cb3 100644 --- a/lightllm/common/speculative/__init__.py +++ b/lightllm/common/speculative/__init__.py @@ -1,21 +1,27 @@ from .config import ( + BlockDraftLayout, SpeculativeConfig, - get_dspark_family_block_size, + get_block_draft_layout, + normalize_speculative_draft_config, is_dspark_draft_config, is_eagle3_draft_config, is_gemma4_dspark_draft_config, is_qwen3_dflash_draft_config, + is_qwen3_5_dflash_draft_config, is_qwen3_dspark_draft_config, validate_dspark_family_draft_config, ) __all__ = [ + "BlockDraftLayout", "SpeculativeConfig", - "get_dspark_family_block_size", + "get_block_draft_layout", + "normalize_speculative_draft_config", "is_dspark_draft_config", "is_eagle3_draft_config", "is_gemma4_dspark_draft_config", "is_qwen3_dflash_draft_config", + "is_qwen3_5_dflash_draft_config", "is_qwen3_dspark_draft_config", "validate_dspark_family_draft_config", ] diff --git a/lightllm/common/speculative/config.py b/lightllm/common/speculative/config.py index 9c92dde73d..3bd0581af0 100644 --- a/lightllm/common/speculative/config.py +++ b/lightllm/common/speculative/config.py @@ -1,5 +1,5 @@ from dataclasses import dataclass -from typing import Any, Mapping, Optional +from typing import Any, Mapping, MutableMapping, Optional VANILLA_SPEC_MODES = frozenset({"vanilla_with_att", "vanilla_no_att", "qwen3next_vanilla"}) @@ -11,10 +11,47 @@ NO_ATTENTION_SPEC_MODES = frozenset({"vanilla_no_att", "eagle_no_att", "qwen3next_vanilla", "qwen3next_eagle"}) TARGET_HIDDEN_SPEC_MODES = frozenset({"eagle3", "dspark", "dflash"}) QWEN3_DFLASH_ARCHITECTURES = frozenset({"Qwen3DFlashModel", "Qwen3DSparkModel"}) +QWEN3_5_DFLASH_ARCHITECTURES = frozenset({"Qwen3_5DFlashModel"}) QWEN3_DSPARK_ARCHITECTURES = frozenset({"Qwen3DSparkModel"}) GEMMA4_DSPARK_ARCHITECTURES = frozenset({"Gemma4DSparkModel"}) -DSPARK_FAMILY_ARCHITECTURES = QWEN3_DFLASH_ARCHITECTURES | GEMMA4_DSPARK_ARCHITECTURES +DSPARK_FAMILY_ARCHITECTURES = QWEN3_DFLASH_ARCHITECTURES | QWEN3_5_DFLASH_ARCHITECTURES | GEMMA4_DSPARK_ARCHITECTURES DSPARK_MARKOV_HEAD_TYPES = frozenset({"vanilla", "gated", "rnn"}) +SPECULATIVE_CONFIG_SECTIONS = ("dflash_config", "dspark_config", "draft_config", "speculative_config", "mtp_config") +SPECULATIVE_DRAFT_CONFIG_KEYS = frozenset( + { + "block_size", + "target_layer_ids", + "mask_token_id", + "markov_rank", + "markov_head_type", + "enable_confidence_head", + "confidence_head_with_markov", + } +) + + +@dataclass(frozen=True) +class BlockDraftLayout: + """Runtime layout of a non-causal block draft checkpoint. + + ``query_block_size`` is the number of logits rows emitted for one anchor, + while ``proposal_output_start`` identifies the first row that represents a + draft token. Keeping both values explicit lets serving support checkpoints + whose query block includes a leading bonus row without teaching generic + scheduling or proposer code about a particular model architecture. + """ + + query_block_size: int + proposal_output_start: int + + @property + def draft_step(self) -> int: + return self.query_block_size - self.proposal_output_start + + def resolve_draft_step(self, configured_step: int) -> int: + """Use a positive configured step up to the checkpoint's proposal capacity.""" + configured_step = int(configured_step) + return configured_step if 0 < configured_step <= self.draft_step else self.draft_step @dataclass(frozen=True) @@ -153,6 +190,11 @@ def is_qwen3_dflash_draft_config(config: Mapping[str, Any]) -> bool: return any(architecture in QWEN3_DFLASH_ARCHITECTURES for architecture in architectures) +def is_qwen3_5_dflash_draft_config(config: Mapping[str, Any]) -> bool: + architectures = config.get("architectures", []) + return any(architecture in QWEN3_5_DFLASH_ARCHITECTURES for architecture in architectures) + + def is_qwen3_dspark_draft_config(config: Mapping[str, Any]) -> bool: architectures = config.get("architectures", []) return any(architecture in QWEN3_DSPARK_ARCHITECTURES for architecture in architectures) @@ -163,8 +205,27 @@ def is_gemma4_dspark_draft_config(config: Mapping[str, Any]) -> bool: return any(architecture in GEMMA4_DSPARK_ARCHITECTURES for architecture in architectures) +def normalize_speculative_draft_config(config: MutableMapping[str, Any]) -> MutableMapping[str, Any]: + """Normalize supported speculative draft config layouts in place. + + LightLLM model code generally normalizes nested checkpoint config once at + load time, then downstream code reads a flat `network_config`. This helper + applies the same pattern to speculative draft checkpoints whose shared + fields may live under sections such as `dflash_config`. + """ + + for section in SPECULATIVE_CONFIG_SECTIONS: + nested_config = config.get(section) + if not isinstance(nested_config, Mapping): + continue + for key in SPECULATIVE_DRAFT_CONFIG_KEYS: + if key not in config and key in nested_config: + config[key] = nested_config[key] + return config + + def validate_dspark_family_draft_config( - config: Mapping[str, Any], + config: MutableMapping[str, Any], *, require_confidence_head: bool = False, ) -> None: @@ -172,6 +233,8 @@ def validate_dspark_family_draft_config( assert is_dspark_draft_config(config), f"unsupported DFlash/DSpark architecture: {config.get('architectures')}" + normalize_speculative_draft_config(config) + block_size = int(config.get("block_size", 0)) assert block_size > 0, "DFlash/DSpark draft config must provide positive block_size" @@ -214,13 +277,34 @@ def validate_dspark_family_draft_config( return -def get_dspark_family_block_size( - config: Mapping[str, Any], +def get_block_draft_layout( + config: MutableMapping[str, Any], *, + mode: str, require_confidence_head: bool = False, -) -> int: +) -> BlockDraftLayout: + """Resolve the generic query/proposal layout of a block draft checkpoint. + + Proposal rows start at zero by default. Architectures with a different + upstream block contract are registered explicitly. + """ + + assert mode in BLOCK_SPEC_MODES, f"block draft layout is not defined for mode {mode!r}" validate_dspark_family_draft_config( config, require_confidence_head=require_confidence_head, ) - return int(config["block_size"]) + + query_block_size = int(config["block_size"]) + # The Z-Lab Qwen3.5 DFlash checkpoint defines block_size as the full + # query block: row 0 is the accepted/bonus query and proposals start at 1. + proposal_output_start = 1 if is_qwen3_5_dflash_draft_config(config) else 0 + + assert 0 <= proposal_output_start < query_block_size, ( + "block draft proposal_output_start must be within the query block: " + f"start={proposal_output_start}, block_size={query_block_size}" + ) + return BlockDraftLayout( + query_block_size=query_block_size, + proposal_output_start=proposal_output_start, + ) diff --git a/lightllm/models/__init__.py b/lightllm/models/__init__.py index d56b17608a..7a9aa3cd8f 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -64,6 +64,7 @@ ), "Qwen3_5TpPartModel": ("lightllm.models.qwen3_5.model", "Qwen3_5TpPartModel"), "Qwen3_5MOETpPartModel": ("lightllm.models.qwen3_5_moe.model", "Qwen3_5MOETpPartModel"), + "Qwen3_5DFlashModel": ("lightllm.models.qwen3_5_dflash.model", "Qwen3_5DFlashModel"), } _MODEL_TYPE_REGISTRY_MODULES = { diff --git a/lightllm/models/qwen3_5_dflash/__init__.py b/lightllm/models/qwen3_5_dflash/__init__.py new file mode 100644 index 0000000000..f09bbdb69c --- /dev/null +++ b/lightllm/models/qwen3_5_dflash/__init__.py @@ -0,0 +1,3 @@ +from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel + +__all__ = ["Qwen3_5DFlashModel"] diff --git a/lightllm/models/qwen3_5_dflash/layer_weights/__init__.py b/lightllm/models/qwen3_5_dflash/layer_weights/__init__.py new file mode 100644 index 0000000000..4a889a89cb --- /dev/null +++ b/lightllm/models/qwen3_5_dflash/layer_weights/__init__.py @@ -0,0 +1,5 @@ +from lightllm.models.qwen3_5_dflash.layer_weights.pre_and_post_layer_weight import ( + Qwen35DFlashPreAndPostLayerWeight, +) + +__all__ = ["Qwen35DFlashPreAndPostLayerWeight"] diff --git a/lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py new file mode 100644 index 0000000000..b22eaeb89d --- /dev/null +++ b/lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,47 @@ +from lightllm.common.basemodel import PreAndPostLayerWeight +from lightllm.common.basemodel.layer_weights.meta_weights import ( + EmbeddingWeight, + LMHeadWeight, + RMSNormWeight, + ROWMMWeight, +) +from lightllm.common.quantization import Quantcfg + + +class Qwen35DFlashPreAndPostLayerWeight(PreAndPostLayerWeight): + """Qwen3.5 DFlash weights outside the draft decoder stack. + + Qwen3.5 DFlash checkpoints only store the DFlash projection and draft + output norm. Token embedding and LM head weights are shared from the target + model by Qwen3_5DFlashModel. + """ + + def __init__(self, data_type, network_config, quant_cfg: Quantcfg): + super().__init__(data_type, network_config) + self.quant_cfg = quant_cfg + + hidden_size = network_config["hidden_size"] + target_layer_num = len(network_config["target_layer_ids"]) + + self.wte_weight_: EmbeddingWeight = None + self.lm_head_weight_: LMHeadWeight = None + self.fc_weight_ = ROWMMWeight( + in_dim=hidden_size * target_layer_num, + out_dims=[hidden_size], + weight_names="fc.weight", + data_type=self.data_type_, + quant_method=self.quant_cfg.get_quant_method(0, "fc"), + tp_rank=0, + tp_world_size=1, + ) + self.hidden_norm_weight_ = RMSNormWeight( + dim=hidden_size, + weight_name="hidden_norm.weight", + data_type=self.data_type_, + ) + self.final_norm_weight_ = RMSNormWeight( + dim=hidden_size, + weight_name="norm.weight", + data_type=self.data_type_, + ) + return diff --git a/lightllm/models/qwen3_5_dflash/model.py b/lightllm/models/qwen3_5_dflash/model.py new file mode 100644 index 0000000000..571fc1f1e7 --- /dev/null +++ b/lightllm/models/qwen3_5_dflash/model.py @@ -0,0 +1,135 @@ +from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager +from lightllm.common.kv_cache_mem_manager.operator import NormalMemOperator +from lightllm.common.speculative.config import normalize_speculative_draft_config +from lightllm.distributed.communication_op import dist_group_manager +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel +from lightllm.models.qwen3_5_dflash.layer_weights.pre_and_post_layer_weight import ( + Qwen35DFlashPreAndPostLayerWeight, +) + + +class Qwen35DFlashMemOperator(NormalMemOperator): + """Map global draft layer indexes onto the draft-owned KV cache. + + Qwen3.5 target models use a Qwen3Next mixed linear/full-attention memory + manager, which is not compatible with ordinary DFlash KV writes. The + Qwen3 DFlash layer infer code is reused here and still passes global layer + indexes, so this operator converts those indexes back to the local + 0..draft_layer_num-1 range used by the independent draft MemoryManager. + """ + + def _local_layer_index(self, layer_index: int) -> int: + return int(layer_index) - int(self.mem_manager.layer_index_offset) + + def copy_kv_to_mem_manager(self, layer_index: int, mem_index, kv): + return super().copy_kv_to_mem_manager(self._local_layer_index(layer_index), mem_index, kv) + + +class Qwen35DFlashMemoryManager(MemoryManager): + """Ordinary KV cache for the Qwen3.5 DFlash draft model. + + Unlike Qwen3 DFlash, this draft cannot share the target model memory + manager because the Qwen3.5 target cache contains Qwen3Next linear-attention + state. The layer_index_offset records where the draft layers sit in the + global model layer list, while the underlying cache stores only draft-local + layer rows. + """ + + operator_class = Qwen35DFlashMemOperator + + def __init__(self, *args, layer_index_offset: int = 0, **kwargs): + self.layer_index_offset = int(layer_index_offset) + super().__init__(*args, **kwargs) + + def get_att_input_params(self, layer_index: int): + return super().get_att_input_params(int(layer_index) - self.layer_index_offset) + + +class Qwen3_5DFlashModel(Qwen3DFlashModel): + """Qwen3.5 DFlash draft model. + + The checkpoint stores DFlash parameters under `dflash_config` and omits + token embedding / LM head weights. Those two weights are shared from the + already-loaded target model. Its query/proposal layout is resolved from the + checkpoint schema by the generic speculative runtime. + """ + + pre_and_post_weight_class = Qwen35DFlashPreAndPostLayerWeight + + def __init__(self, kvargs: dict): + super().__init__(kvargs) + # The request manager is shared with the target, while this draft owns + # a separate KV cache. TpPartBaseModel wires every request manager to + # the model's memory manager during construction, which would make + # request allocation bypass the target/radix-cache allocator. Restore + # the shared request manager's allocator after draft initialization. + self._restore_main_mem_manager() + return + + def _restore_main_mem_manager(self): + assert self.req_manager is self.main_model.req_manager + self.req_manager.mem_manager = self.main_model.mem_manager + return + + def _init_config(self): + super()._init_config() + self._normalize_dflash_config() + return + + def _normalize_dflash_config(self): + normalize_speculative_draft_config(self.config) + + rope_parameters = self.config.get("rope_parameters") + if isinstance(rope_parameters, dict): + if "rope_theta" in rope_parameters and "rope_theta" not in self.config: + self.config["rope_theta"] = rope_parameters["rope_theta"] + if "partial_rotary_factor" in rope_parameters and "partial_rotary_factor" not in self.config: + self.config["partial_rotary_factor"] = rope_parameters["partial_rotary_factor"] + if "rope_scaling" not in self.config: + self.config["rope_scaling"] = rope_parameters + return + + def _init_custom(self): + # The Qwen3.5 target uses mrope/partial rotary with a different head + # shape from the ordinary DFlash draft. Build draft-owned rotary caches + # from the draft config instead of reusing main_model._cos/_sin. + LlamaTpPartModel._init_custom(self) + self.dist_group = dist_group_manager.get_default_group() + self.block_size = int(self.config["block_size"]) + self.mask_token_id = int(self.config["mask_token_id"]) + return + + def _init_mem_manager(self): + head_dim = self.config.get("head_dim") + if head_dim is None: + head_dim = self.config["hidden_size"] // self.config["num_attention_heads"] + num_kv_heads = max(1, self.config["num_key_value_heads"] // self.tp_world_size_) + self.draft_layer_start = len(self.main_model.layers_infer) + self.draft_layer_start += sum( + len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models + ) + self.mem_manager = Qwen35DFlashMemoryManager( + size=self.main_model.mem_manager.size, + dtype=self.data_type, + head_num=num_kv_heads, + head_dim=head_dim, + layer_num=self.config["n_layer"], + mem_fraction=self.mem_fraction, + layer_index_offset=self.draft_layer_start, + ) + register_draft_manager = getattr( + self.main_model.mem_manager, "register_dflash_draft_mem_manager", None + ) + assert callable(register_draft_manager), "Qwen3_5DFlashModel requires a Qwen3Next target" + register_draft_manager( + self.mem_manager, + global_kv_heads=int(self.config["num_key_value_heads"]), + ) + return + + def _init_weights(self, start_layer_index=None): + super()._init_weights(start_layer_index=start_layer_index) + self.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ + self.pre_post_weight.lm_head_weight_ = self.main_model.pre_post_weight.lm_head_weight_ + return diff --git a/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py index 54680af804..31a6fba7a7 100644 --- a/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py @@ -1,11 +1,21 @@ import torch +from lightllm.common.basemodel.attention import AttControl from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo from lightllm.models.qwen3_dflash.layer_weights.transformer_layer_weight import Qwen3DFlashTransformerLayerWeight +def get_draft_layer_type(layer_num, network_config): + layer_types = network_config.get("layer_types", []) + draft_layer_start = int(network_config.get("_draft_layer_start", 0)) + local_layer_num = int(layer_num) - draft_layer_start + if 0 <= local_layer_num < len(layer_types): + return layer_types[local_layer_num] + return "full_attention" + + class Qwen3DFlashTransformerLayerInfer(LlamaTransformerLayerInfer): """DFlash layer inference. @@ -18,6 +28,11 @@ def __init__(self, layer_num, network_config): super().__init__(layer_num, network_config) self.head_dim_ = network_config["head_dim"] self.block_size_ = int(network_config["block_size"]) + layer_type = get_draft_layer_type(layer_num, network_config) + sliding_window = int(network_config.get("sliding_window", 0) or 0) + self.use_sliding_window_ = bool(network_config.get("use_sliding_window", False)) + self.use_sliding_window_ = self.use_sliding_window_ and layer_type == "sliding_attention" and sliding_window > 0 + self.sliding_window_ = sliding_window return def context_forward( @@ -111,3 +126,24 @@ def _get_qkv(self, input, infer_state: Qwen3DFlashInferStateInfo, layer_weight: self.head_dim_, ) return q, cache_kv + + def _token_attention_kernel( + self, + q: torch.Tensor, + infer_state: Qwen3DFlashInferStateInfo, + layer_weight: Qwen3DFlashTransformerLayerWeight, + ) -> torch.Tensor: + _k, _v = infer_state.mem_manager.get_att_input_params(layer_index=self.layer_num_) + _q = q.view(-1, self.tp_q_head_num_, self.head_dim_) + if self.use_sliding_window_: + att_control = AttControl(use_sliding_window=True, sliding_window=(self.sliding_window_ - 1, 0)) + else: + att_control = AttControl() + o_tensor = infer_state.decode_att_state.decode_att( + q=_q, + k=_k, + v=_v, + att_control=att_control, + alloc_func=self.alloc_tensor, + ) + return o_tensor.view(q.shape) diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py index cae7990bb4..ea01f18251 100644 --- a/lightllm/models/qwen3_dflash/model.py +++ b/lightllm/models/qwen3_dflash/model.py @@ -109,6 +109,7 @@ def _init_infer_layer(self, start_layer_index=None): self.draft_layer_start += sum( len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models ) + self.config["_draft_layer_start"] = self.draft_layer_start super()._init_infer_layer(start_layer_index=self.draft_layer_start) return diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 5c4879060f..c19096c7f5 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -28,7 +28,10 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args -from lightllm.common.speculative import SpeculativeConfig, get_dspark_family_block_size +from lightllm.common.speculative import ( + SpeculativeConfig, + get_block_draft_layout, +) logger = init_logger(__name__) @@ -41,20 +44,23 @@ def normalize_block_mtp_step_from_first_draft_config( assert args.mtp_draft_model_dir is not None and len(args.mtp_draft_model_dir) > 0 mtp_model_cfg, _ = PretrainedConfig.get_config_dict(args.mtp_draft_model_dir[0]) - block_size = get_dspark_family_block_size( + layout = get_block_draft_layout( mtp_model_cfg, + mode=spec_config.mode, require_confidence_head=spec_config.is_dspark, ) configured_step = int(args.mtp_step) - if configured_step not in (0, block_size): + draft_step = layout.resolve_draft_step(configured_step) + if configured_step not in (0, draft_step): logger.warning( - "Overriding mtp_step=%s with block draft config block_size=%s for %s mode", + "Overriding mtp_step=%s with draft_step=%s from block_size=%s for %s mode", configured_step, - block_size, + draft_step, + layout.query_block_size, spec_config.mode, ) - args.mtp_step = block_size - spec_config = replace(spec_config, step=block_size) + args.mtp_step = draft_step + spec_config = replace(spec_config, step=draft_step) spec_config.validate() return spec_config diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index d0f32419d6..524d1a97b4 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -184,7 +184,7 @@ def __post_init__(self): error_info = "start_kv_index must >=0 and end_kv_index > start_kv_index" logger.error(error_info) raise ValueError(error_info) - if self.page_kind == "kv": + if self.page_kind in ("kv", "draft_kv"): assert len(self.mem_indexes) == (self.end_kv_index - self.start_kv_index) elif self.page_kind == "linear_att_state": assert self.start_kv_index == self.end_kv_index diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 8f9273336c..7c350bcc54 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -436,6 +436,7 @@ def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: L copy_linear_att_state_to_kv_buffer( b_req_idx=b_req_idx, + req_to_mtp_state_index=self.req_manager.req_to_mtp_state_index, big_page_buffer_ids=big_page_buffer_ids, gpu_conv_state=self.req_manager.req_to_conv_state.buffer, gpu_ssm_state=self.req_manager.req_to_ssm_state.buffer, 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 d9e9b160ed..961044b757 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -40,11 +40,12 @@ from lightllm.distributed import dist_group_manager from lightllm.common.speculative import ( SpeculativeConfig, - get_dspark_family_block_size, + get_block_draft_layout, is_dspark_draft_config, is_eagle3_draft_config, is_gemma4_dspark_draft_config, is_qwen3_dflash_draft_config, + is_qwen3_5_dflash_draft_config, is_qwen3_dspark_draft_config, ) from lightllm.server.router.model_infer.speculative import build_spec_runtime @@ -185,6 +186,7 @@ def init_model(self, kvargs): self.model: TpPartBaseModel = self.model # for easy typing set_random_seed(2147483647) self.is_linear_att_mixed_model = isinstance(self.model.req_manager, ReqManagerForMamba) + self._validate_linear_att_spec_support() if self.is_linear_att_mixed_model: self.linear_att_cache_manager = LinearAttCacheManager( @@ -250,12 +252,6 @@ def init_model(self, kvargs): [rank for rank in range(self.global_world_size)], backend="nccl" ) - if self.args.run_mode in ["prefill", "decode"] or self.args.enable_dp_prompt_cache_fetch: - # 如果存在需要跨进程使用mem manger的特性,则将mem manager写入到 shm中,方便 - # 读取 - self.model.mem_manager.write_to_shm(req_manager=self.model.req_manager) - dist.barrier(group=self.node_nccl_group) - # 同一 DP 组内只需主 rank 初始化真实的 capture buffer 并执行后续相关操作; # 非主 rank 不需要分配 buffer,避免重复占用内存。 if self.is_master_in_dp: @@ -271,9 +267,6 @@ def init_model(self, kvargs): self.init_custom() - if self.args.enable_dp_prompt_cache_fetch: - self.init_dp_kv_shared() - self.shm_reqs_io_buffer = ShmObjsIOBuffer() # 只会在 pd pd 模式下才会使用,用于上传分块传输任务是否成功。 self.shm_pd_trans_io_buffer = ShmObjsIOBuffer(tail_str="pd") @@ -286,6 +279,16 @@ def init_model(self, kvargs): self.spec_adapter = build_spec_runtime(self) self._attach_spec_adapter() + if self.args.run_mode in ["prefill", "decode"] or self.args.enable_dp_prompt_cache_fetch: + # Draft models must be initialized before this snapshot. Qwen3.5 + # DFlash attaches its independent draft KV manager to the target + # manager so PD transfer workers can deserialize both buffers. + self.model.mem_manager.write_to_shm(req_manager=self.model.req_manager) + dist.barrier(group=self.node_nccl_group) + + if self.args.enable_dp_prompt_cache_fetch: + self.init_dp_kv_shared() + if self.args.enable_cpu_cache: self.multi_level_cache_module = MultiLevelKvCacheModule(self) @@ -372,7 +375,6 @@ def init_mtp_draft_model(self, main_kvargs: dict): for i in range(num_mtp_modules): mtp_model_cfg, _ = PretrainedConfig.get_config_dict(mtp_draft_model_dirs[i]) - self._normalize_block_mtp_step_from_config(mtp_model_cfg) spec_config = self.spec_config model_type = mtp_model_cfg.get("model_type", "") mtp_model_kvargs = { @@ -427,6 +429,10 @@ def init_mtp_draft_model(self, main_kvargs: dict): from lightllm.models.qwen3_eagle.model import Qwen3EagleModel self.draft_models.append(Qwen3EagleModel(mtp_model_kvargs)) + elif spec_config.is_dflash and is_qwen3_5_dflash_draft_config(mtp_model_cfg): + from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel + + self.draft_models.append(Qwen3_5DFlashModel(mtp_model_kvargs)) elif spec_config.is_dflash and is_qwen3_dflash_draft_config(mtp_model_cfg): from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel @@ -449,21 +455,46 @@ def _normalize_block_mtp_step_from_config(self, mtp_model_cfg: dict) -> None: if not self.spec_config.uses_block_draft_model: return - block_size = get_dspark_family_block_size( + layout = get_block_draft_layout( mtp_model_cfg, + mode=self.spec_config.mode, require_confidence_head=self.spec_config.is_dspark, ) + self.block_draft_layout = layout configured_step = int(getattr(self.args, "mtp_step", 0)) - if configured_step not in (0, block_size): + draft_step = layout.resolve_draft_step(configured_step) + if configured_step not in (0, draft_step): self.logger.warning( - "Overriding mtp_step=%s with block draft config block_size=%s for %s mode", + "Overriding mtp_step=%s with draft_step=%s from block_size=%s for %s mode", configured_step, - block_size, + draft_step, + layout.query_block_size, self.spec_config.mode, ) - self.args.mtp_step = block_size - self.mtp_step = block_size - self.spec_config = replace(self.spec_config, step=block_size) + self.args.mtp_step = draft_step + self.mtp_step = draft_step + self.spec_config = replace(self.spec_config, step=draft_step) + return + + def _validate_linear_att_spec_support(self) -> None: + """Restrict the new DFlash combination without narrowing existing LightSpec modes.""" + + if not self.spec_config.enabled or not self.spec_config.is_dflash: + return + + mtp_draft_model_dirs = self.args.mtp_draft_model_dir + if isinstance(mtp_draft_model_dirs, str): + mtp_draft_model_dirs = [mtp_draft_model_dirs] + assert mtp_draft_model_dirs is not None and len(mtp_draft_model_dirs) == 1 + mtp_model_cfg, _ = PretrainedConfig.get_config_dict(mtp_draft_model_dirs[0]) + is_qwen35_dflash = is_qwen3_5_dflash_draft_config(mtp_model_cfg) + if self.is_linear_att_mixed_model: + assert is_qwen35_dflash, ( + "linear-attention mixed targets require a Qwen3_5DFlashModel draft checkpoint, " + f"got architectures={mtp_model_cfg.get('architectures')}" + ) + else: + assert not is_qwen35_dflash, "Qwen3_5DFlashModel requires a Qwen3Next target" return def _normalize_block_mtp_step_from_first_draft_config(self) -> None: diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index f9dc6ee60f..1a20fc4f82 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -144,6 +144,15 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): kv_end_index=end_index, group=group, ) + if getattr(self.model.mem_manager, "has_separate_dflash_draft_kv", False): + self._create_pd_trans_task( + req_obj=req_obj, + mem_indexes=page_mem_indexes.tolist(), + kv_start_index=start_index, + kv_end_index=end_index, + group=group, + page_kind="draft_kv", + ) # update req_obj.pd_trans_kv_start_index += cur_page_size @@ -194,7 +203,7 @@ def _create_pd_trans_task( # only self.is_master_in_dp will be used. self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size - if page_kind == "kv": + if page_kind in ("kv", "draft_kv"): req_idx = None elif page_kind == "linear_att_state": req_idx = req_obj.req_idx diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index 2a501f509b..4fbb1eaf89 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -63,13 +63,25 @@ def _prefill_chuncked_handle_func( cur_page_size = min(page_size, req_obj.cur_kv_len - req_obj.pd_trans_kv_start_index) # 生成页面传输任务, 放入kv move manager 的处理队列中 if cur_page_size == page_size or prefill_finished: - trans_task = self._create_pd_trans_task( - req_obj=req_obj, - kv_start_index=req_obj.pd_trans_kv_start_index, - kv_end_index=req_obj.pd_trans_kv_start_index + cur_page_size, + start_index = req_obj.pd_trans_kv_start_index + end_index = start_index + cur_page_size + trans_task_list.append( + self._create_pd_trans_task( + req_obj=req_obj, + kv_start_index=start_index, + kv_end_index=end_index, + ) ) + if getattr(self.model.mem_manager, "has_separate_dflash_draft_kv", False): + trans_task_list.append( + self._create_pd_trans_task( + req_obj=req_obj, + kv_start_index=start_index, + kv_end_index=end_index, + page_kind="draft_kv", + ) + ) req_obj.pd_trans_kv_start_index += cur_page_size - trans_task_list.append(trans_task) else: break @@ -106,7 +118,7 @@ def _create_pd_trans_task( self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size pd_decode_node_info = req_obj.sampling_param.pd_decode_node - if page_kind == "kv": + if page_kind in ("kv", "draft_kv"): mem_indexes = ( self.model.req_manager.req_to_token_indexs[req_obj.req_idx, kv_start_index:kv_end_index] .detach() diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index 82b0c5c431..4936f89d65 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -5,6 +5,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.speculative import BlockDraftLayout from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal @@ -63,7 +64,8 @@ def propose_next( num_reqs = int(b_req_mtp_start_loc.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - assert block_size >= draft_step + layout = self.backend.block_draft_layout + assert block_size == layout.query_block_size assert verify_result.accept_len.shape[0] == num_reqs token_ids = next_token_ids.new_full( (next_token_ids.shape[0], draft_step + 1), @@ -78,7 +80,10 @@ def propose_next( draft_probs=None, ) - self.extend_draft_kv_cache(main_model_input=main_model_input) + self.extend_draft_kv_cache( + main_model_input=main_model_input, + accepted_index=verify_result.accepted_index, + ) # DFlash only drafts from the accepted tail row of each request. Unlike # MTP, one anchor row expands to a whole non-causal block. @@ -97,26 +102,60 @@ def propose_next( flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) assert flat_token_ids.numel() == num_reqs * block_size block_token_ids = flat_token_ids.reshape(num_reqs, block_size) - token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] + token_ids[selected_rows, 1:] = self.select_draft_token_ids( + block_token_ids=block_token_ids, + draft_step=draft_step, + layout=layout, + ) return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=draft_mem_indexes_cpu, draft_probs=None, ) - def extend_draft_kv_cache(self, *, main_model_input: ModelInput) -> None: + @staticmethod + def select_draft_token_ids( + *, + block_token_ids: torch.Tensor, + draft_step: int, + layout: BlockDraftLayout, + ) -> torch.Tensor: + output_start = layout.proposal_output_start + output_end = output_start + int(draft_step) + assert 0 <= output_start <= output_end <= block_token_ids.shape[1] + return block_token_ids[:, output_start:output_end] + + def extend_draft_kv_cache(self, *, main_model_input: ModelInput, accepted_index: torch.Tensor) -> None: target_hidden = self.runtime.get_hidden() draft_model = self.backend.draft_models[0] + accepted_rows = torch.nonzero(accepted_index.to(torch.bool), as_tuple=False).flatten().to(torch.long) + if accepted_rows.numel() == 0: + return + target_hidden = target_hidden.index_select(0, accepted_rows) + batch_size = int(target_hidden.shape[0]) draft_kv_input = copy.copy(main_model_input) draft_kv_input.batch_size = batch_size draft_kv_input.total_token_num = batch_size + draft_kv_input.multimodal_params = [{"images": [], "audios": []} for _ in range(batch_size)] + # This hidden-commit prefill path does not consume token ids, but + # InferState uses input_ids.shape[0] to build position ids. Keep it + # aligned with the accepted hidden rows. + draft_kv_input.input_ids = torch.empty( + (batch_size,), + dtype=torch.int64, + device=target_hidden.device, + ) draft_kv_input.max_q_seq_len = 1 draft_kv_input.prefix_total_token_num = 0 draft_kv_input.is_prefill = True - # Each expanded MTP row writes one target-hidden KV slot for the same request. - draft_kv_input.b_ready_cache_len = main_model_input.b_seq_len - 1 + # Each accepted MTP row writes one target-hidden KV slot for the same request. + draft_kv_input.b_req_idx = main_model_input.b_req_idx.index_select(0, accepted_rows).contiguous() + draft_kv_input.b_mtp_index = main_model_input.b_mtp_index.index_select(0, accepted_rows).contiguous() + draft_kv_input.b_seq_len = main_model_input.b_seq_len.index_select(0, accepted_rows).contiguous() + draft_kv_input.mem_indexes = main_model_input.mem_indexes.index_select(0, accepted_rows).contiguous() + draft_kv_input.b_ready_cache_len = draft_kv_input.b_seq_len - 1 draft_kv_input.b_prefill_start_loc = torch.arange( batch_size, dtype=torch.int32, diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/speculative/proposers/dspark.py index 4340e2754d..aca931f759 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dspark.py @@ -36,7 +36,8 @@ def propose_next( verify_row_count = next_token_ids.shape[0] draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - assert block_size >= draft_step + layout = self.backend.block_draft_layout + assert block_size == layout.query_block_size assert verify_result.accept_len.shape[0] == num_reqs proposal_token_ids = next_token_ids.new_full( @@ -54,7 +55,10 @@ def propose_next( schedule_probs=schedule_probs, ) - self.extend_draft_kv_cache(main_model_input=main_model_input) + self.extend_draft_kv_cache( + main_model_input=main_model_input, + accepted_index=verify_result.accepted_index, + ) selected_rows = self.select_accepted_tail_rows( b_req_mtp_start_loc=b_req_mtp_start_loc, accept_len=verify_result.accept_len, @@ -77,7 +81,11 @@ def propose_next( flat_token_ids.numel() == expected_block_rows ), f"draft token rows must be {expected_block_rows}, got {flat_token_ids.numel()}" block_token_ids = flat_token_ids.reshape(num_reqs, block_size) - proposal_token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] + proposal_token_ids[selected_rows, 1:] = self.select_draft_token_ids( + block_token_ids=block_token_ids, + draft_step=draft_step, + layout=layout, + ) draft_probs = None if self.enable_dynamic_mtp: @@ -89,11 +97,13 @@ def propose_next( confidence_logits.shape[0] == num_reqs ), f"confidence logits rows must be {num_reqs}, got {confidence_logits.shape[0]}" assert ( - confidence_logits.shape[1] >= draft_step - ), f"confidence logits columns must cover draft_step={draft_step}, got {confidence_logits.shape[1]}" + confidence_logits.shape[1] >= layout.proposal_output_start + draft_step + ), f"confidence logits columns must cover the proposal layout, got {confidence_logits.shape[1]}" + output_start = layout.proposal_output_start + output_end = output_start + draft_step schedule_probs = self._scatter_step_probs( selected_rows=selected_rows, - probs=confidence_logits[:, :draft_step].sigmoid(), + probs=confidence_logits[:, output_start:output_end].sigmoid(), verify_row_count=verify_row_count, ) diff --git a/lightllm/server/router/model_infer/speculative/runtime.py b/lightllm/server/router/model_infer/speculative/runtime.py index 4869109711..91cf7c51cd 100644 --- a/lightllm/server/router/model_infer/speculative/runtime.py +++ b/lightllm/server/router/model_infer/speculative/runtime.py @@ -7,7 +7,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.speculative.config import SpeculativeConfig +from lightllm.common.speculative.config import SpeculativeConfig, normalize_speculative_draft_config from lightllm.server.router.model_infer.speculative.planner import FixedMTPPlanner, SpecDecodePlan from lightllm.server.router.model_infer.speculative.proposers import build_spec_proposer from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal @@ -745,6 +745,7 @@ def _get_target_layer_ids(self, model) -> List[int]: if draft_model_dir: with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: draft_config = json.load(json_file) + normalize_speculative_draft_config(draft_config) target_layer_ids = draft_config.get("target_layer_ids") if target_layer_ids is not None: self._target_layer_ids = [int(layer_id) for layer_id in target_layer_ids] diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index b16e1612ac..119e21cb7c 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -5,7 +5,7 @@ from easydict import EasyDict from functools import lru_cache from lightllm.utils.log_utils import init_logger -from lightllm.common.speculative import SpeculativeConfig +from lightllm.common.speculative import SpeculativeConfig, is_qwen3_5_dflash_draft_config logger = init_logger(__name__) @@ -279,6 +279,10 @@ def get_added_mtp_kv_layer_num() -> int: return spec_config.draft_model_count with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: draft_config = json.load(json_file) + if spec_config.is_dflash and is_qwen3_5_dflash_draft_config(draft_config): + # Qwen3.5 DFlash owns a separate KV manager; do not reserve duplicate + # draft layers in the mixed target manager. + return 0 return int(draft_config.get("num_hidden_layers", draft_config.get("n_layer", spec_config.draft_model_count))) return spec_config.draft_model_count diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index 0b330b7602..43864aa56b 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -16,7 +16,7 @@ get_added_mtp_kv_layer_num, ) from lightllm.utils.log_utils import init_logger -from lightllm.common.speculative import SpeculativeConfig +from lightllm.common.speculative import SpeculativeConfig, is_qwen3_5_dflash_draft_config from lightllm.utils.config_utils import get_num_key_value_heads, get_head_dim, get_layer_num, is_linear_att_mixed_model from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class from lightllm.common.kv_cache_mem_manager import ( @@ -59,6 +59,35 @@ def compute_token_list_hash(tokens: List[int], cpu_cache_token_page_size: int) - return chunks_hash_value +def _get_qwen35_dflash_cpu_cache_bytes(args) -> int: + """Size of the global draft KV appended to each Qwen3Next CPU cache page.""" + + draft_model_dirs = args.mtp_draft_model_dir + if isinstance(draft_model_dirs, str): + draft_model_dirs = [draft_model_dirs] + assert draft_model_dirs and len(draft_model_dirs) == 1 + + from transformers.configuration_utils import PretrainedConfig + + draft_config, _ = PretrainedConfig.get_config_dict(draft_model_dirs[0]) + assert is_qwen3_5_dflash_draft_config(draft_config), ( + "linear-attention MTP CPU cache only supports Qwen3_5DFlashModel, " + f"got architectures={draft_config.get('architectures')}" + ) + page_size = args.cpu_cache_token_page_size + draft_bytes = ( + page_size + * get_layer_num(draft_model_dirs[0]) + * 2 + * get_num_key_value_heads(draft_model_dirs[0]) + * get_head_dim(draft_model_dirs[0]) + * get_llm_data_type().itemsize + ) + # The following typed draft view requires its byte offset and extent to be + # naturally aligned. Existing mixed-model pages use the same alignment. + return triton.cdiv(draft_bytes, 16) * 16 + + @lru_cache(maxsize=None) def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": args = get_env_start_args() @@ -121,11 +150,11 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": spec_config = SpeculativeConfig.from_args(args) if spec_config.enabled: - # TODO 可能会存在不同mtp模式的精度问题 - if not is_linear_att_mixed_model(args.model_dir): - # 对于非 linear att 混合模型,需要额外增加 mtp 的 kv 层数, - # 对于 linear att 混合模型,如qwen 3.5 mtp,已经将 kv 数据 - # 打包成一个块了,所以不需要额外增加,其 layer_num 一直都保持为 1 + if mem_manager_class is Qwen3NextMemManager: + if spec_config.is_dflash: + cpu_cache_meta.head_dim += _get_qwen35_dflash_cpu_cache_bytes(args) + else: + # TODO 可能会存在不同mtp模式的精度问题 cpu_cache_meta.layer_num += get_added_mtp_kv_layer_num() cpu_cache_page_num = int( diff --git a/test/speculative/test_qwen35_dflash_state.py b/test/speculative/test_qwen35_dflash_state.py new file mode 100644 index 0000000000..eed481e683 --- /dev/null +++ b/test/speculative/test_qwen35_dflash_state.py @@ -0,0 +1,378 @@ +from types import MethodType, SimpleNamespace + +import pytest +import torch + +from lightllm.common.speculative import BlockDraftLayout, SpeculativeConfig, get_block_draft_layout +from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel +from lightllm.models.qwen3_dflash.layer_infer.transformer_layer_infer import ( + get_draft_layer_type, +) +from lightllm.common.kv_cache_mem_manager.qwen3next_mem_manager import Qwen3NextMemManager +from lightllm.server import api_start +from lightllm.utils import envs_utils, kv_cache_utils +from lightllm.server.router.model_infer.mode_backend.pd.prefill_node_impl.prefill_impl import ( + PDChunkedPrefillForPrefillNode, +) +from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend +from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer + + +def test_block_layout_skips_leading_bonus_query_prediction(): + block_token_ids = torch.arange(32).reshape(2, 16) + + selected = DFlashProposer.select_draft_token_ids( + block_token_ids=block_token_ids, + draft_step=15, + layout=BlockDraftLayout(query_block_size=16, proposal_output_start=1), + ) + + torch.testing.assert_close(selected, block_token_ids[:, 1:]) + + +def test_block_layout_keeps_first_query_prediction(): + block_token_ids = torch.arange(32).reshape(2, 16) + + selected = DFlashProposer.select_draft_token_ids( + block_token_ids=block_token_ids, + draft_step=16, + layout=BlockDraftLayout(query_block_size=16, proposal_output_start=0), + ) + + torch.testing.assert_close(selected, block_token_ids) + + +def test_block_draft_layout_defaults_to_zero_and_uses_architecture_override(): + draft_fields = { + "block_size": 16, + "target_layer_ids": [1, 10], + "mask_token_id": 248077, + } + nested = { + "architectures": ["Qwen3DFlashModel"], + "dflash_config": draft_fields, + } + flat = {"architectures": ["Qwen3DFlashModel"], **draft_fields, "block_size": 7} + qwen35 = { + "architectures": ["Qwen3_5DFlashModel"], + "dflash_config": draft_fields, + } + + assert get_block_draft_layout( + nested, + mode="dflash", + ) == BlockDraftLayout(query_block_size=16, proposal_output_start=0) + assert get_block_draft_layout( + flat, + mode="dflash", + ) == BlockDraftLayout(query_block_size=7, proposal_output_start=0) + assert get_block_draft_layout(qwen35, mode="dflash") == BlockDraftLayout( + query_block_size=16, proposal_output_start=1 + ) + + +def test_qwen35_dflash_normalizes_block_size_to_fifteen_draft_tokens(): + backend = ModeBackend.__new__(ModeBackend) + backend.args = SimpleNamespace(mtp_step=16) + backend.mtp_step = 16 + backend.spec_config = SpeculativeConfig(mode="dflash", step=16) + backend.logger = SimpleNamespace(warning=lambda *args: None) + config = { + "architectures": ["Qwen3_5DFlashModel"], + "dflash_config": { + "block_size": 16, + "target_layer_ids": [1, 10], + "mask_token_id": 248077, + }, + } + + backend._normalize_block_mtp_step_from_config(config) + + assert backend.args.mtp_step == 15 + assert backend.mtp_step == 15 + assert backend.spec_config.step == 15 + assert backend.block_draft_layout == BlockDraftLayout(query_block_size=16, proposal_output_start=1) + + +def test_qwen35_dflash_respects_shorter_configured_draft_step(): + backend = ModeBackend.__new__(ModeBackend) + backend.args = SimpleNamespace(mtp_step=9) + backend.mtp_step = 9 + backend.spec_config = SpeculativeConfig(mode="dflash", step=9) + backend.logger = SimpleNamespace(warning=lambda *args: None) + config = { + "architectures": ["Qwen3_5DFlashModel"], + "dflash_config": { + "block_size": 16, + "target_layer_ids": [1, 10], + "mask_token_id": 248077, + }, + } + + backend._normalize_block_mtp_step_from_config(config) + + assert backend.args.mtp_step == 9 + assert backend.mtp_step == 9 + assert backend.spec_config.step == 9 + + +def test_qwen35_dflash_normalizes_global_start_args_to_fifteen(monkeypatch): + config = { + "architectures": ["Qwen3_5DFlashModel"], + "dflash_config": { + "block_size": 16, + "target_layer_ids": [1, 10], + "mask_token_id": 248077, + }, + } + monkeypatch.setattr( + api_start.PretrainedConfig, + "get_config_dict", + lambda *_args, **_kwargs: (config, {}), + ) + args = SimpleNamespace(mtp_step=16, mtp_draft_model_dir=["unused"]) + + spec_config = api_start.normalize_block_mtp_step_from_first_draft_config( + args, + SpeculativeConfig(mode="dflash", step=16), + ) + + assert args.mtp_step == 15 + assert spec_config.step == 15 + + +def test_qwen35_dflash_global_normalization_respects_shorter_step(monkeypatch): + config = { + "architectures": ["Qwen3_5DFlashModel"], + "dflash_config": { + "block_size": 16, + "target_layer_ids": [1, 10], + "mask_token_id": 248077, + }, + } + monkeypatch.setattr( + api_start.PretrainedConfig, + "get_config_dict", + lambda *_args, **_kwargs: (config, {}), + ) + args = SimpleNamespace(mtp_step=9, mtp_draft_model_dir=["unused"]) + + spec_config = api_start.normalize_block_mtp_step_from_first_draft_config( + args, + SpeculativeConfig(mode="dflash", step=9), + ) + + assert args.mtp_step == 9 + assert spec_config.step == 9 + + +def test_qwen35_dflash_support_scope_preserves_existing_lightspec_modes(monkeypatch): + backend = ModeBackend.__new__(ModeBackend) + backend.is_linear_att_mixed_model = True + backend.args = SimpleNamespace(mtp_draft_model_dir=["unused"]) + + backend.spec_config = SpeculativeConfig(mode="qwen3next_eagle", step=2, dynamic_verify=True) + backend._validate_linear_att_spec_support() + + backend.spec_config = SpeculativeConfig(mode="dflash", step=6) + monkeypatch.setattr( + "lightllm.server.router.model_infer.mode_backend.base_backend.PretrainedConfig.get_config_dict", + lambda *_args, **_kwargs: ({"architectures": ["Qwen3DFlashModel"]}, {}), + ) + with pytest.raises(AssertionError, match="Qwen3_5DFlashModel"): + backend._validate_linear_att_spec_support() + + monkeypatch.setattr( + "lightllm.server.router.model_infer.mode_backend.base_backend.PretrainedConfig.get_config_dict", + lambda *_args, **_kwargs: ({"architectures": ["Qwen3_5DFlashModel"]}, {}), + ) + backend.is_linear_att_mixed_model = False + with pytest.raises(AssertionError, match="requires a Qwen3Next target"): + backend._validate_linear_att_spec_support() + + backend.is_linear_att_mixed_model = True + backend._validate_linear_att_spec_support() + + + +def test_qwen35_dflash_restores_target_allocator_on_shared_req_manager(): + target_mem_manager = object() + draft_mem_manager = object() + req_manager = SimpleNamespace(mem_manager=draft_mem_manager) + model = Qwen3_5DFlashModel.__new__(Qwen3_5DFlashModel) + model.main_model = SimpleNamespace(req_manager=req_manager, mem_manager=target_mem_manager) + model.req_manager = req_manager + model.mem_manager = draft_mem_manager + + model._restore_main_mem_manager() + + assert model.req_manager.mem_manager is target_mem_manager + assert model.mem_manager is draft_mem_manager + + +def test_qwen35_pd_move_page_supports_distinct_target_and_draft_kv_shapes(): + manager = Qwen3NextMemManager.__new__(Qwen3NextMemManager) + manager.size = 16 + manager.dtype = torch.float32 + manager.layer_num = 2 + manager.head_dim = 4 + manager.linear_config = SimpleNamespace(full_att_all_num_kv_heads=3) + manager._pd_dflash_draft_mem_manager = None + manager._pd_dflash_global_kv_heads = None + draft_manager = SimpleNamespace( + size=16, + dtype=torch.float32, + layer_num=5, + head_dim=2, + ) + + manager.register_dflash_draft_mem_manager(draft_manager, global_kv_heads=2) + elements_per_token = max(2 * 2 * 3 * 4, 5 * 2 * 2 * 2) + manager.kv_move_buffer = torch.empty((1, 7, 1, 1, elements_per_token)) + + target_page = manager._get_pd_kv_page(0, "kv") + draft_page = manager._get_pd_kv_page(0, "draft_kv") + assert target_page.shape == (7, 2, 6, 4) + assert draft_page.shape == (7, 5, 4, 2) + assert target_page.is_contiguous() + assert draft_page.is_contiguous() + + +def test_qwen35_pd_prefill_emits_target_and_draft_kv_tasks(): + emitted_page_kinds = [] + + backend = PDChunkedPrefillForPrefillNode.__new__(PDChunkedPrefillForPrefillNode) + backend.args = SimpleNamespace(pd_kv_page_size=4) + backend.model = SimpleNamespace(mem_manager=SimpleNamespace(has_separate_dflash_draft_kv=True)) + backend.is_master_in_dp = False + + def fake_create_task(self, req_obj, kv_start_index, kv_end_index, page_kind="kv"): + del self, req_obj, kv_start_index, kv_end_index + emitted_page_kinds.append(page_kind) + return SimpleNamespace(first_gen_token_id=None, first_gen_token_logprob=None) + + backend._create_pd_trans_task = MethodType(fake_create_task, backend) + req = SimpleNamespace( + cur_kv_len=4, + pd_trans_kv_start_index=0, + shm_req=SimpleNamespace(input_len=4), + ) + + backend._prefill_chuncked_handle_func(req, next_token_id=1, next_token_prob=0.0, output_len=0) + + assert emitted_page_kinds == ["kv", "draft_kv"] + assert req.pd_trans_kv_start_index == 4 + + +def test_dflash_layer_type_uses_draft_local_index_for_global_layer_numbers(): + config = { + "_draft_layer_start": 64, + "layer_types": ["sliding_attention"] * 5 + ["full_attention"], + } + + assert [get_draft_layer_type(i, config) for i in range(64, 70)] == [ + "sliding_attention", + "sliding_attention", + "sliding_attention", + "sliding_attention", + "sliding_attention", + "full_attention", + ] + + + +def test_qwen35_cpu_cache_appends_distinct_draft_kv_region(monkeypatch): + page_num = 2 + page_size = 4 + target_bytes = 32 + manager = Qwen3NextMemManager.__new__(Qwen3NextMemManager) + manager.size = 16 + manager.dtype = torch.float32 + manager.linear_config = SimpleNamespace(get_cpu_cache_big_page_bytes=lambda: target_bytes) + manager._pd_dflash_draft_mem_manager = None + manager._pd_dflash_global_kv_heads = None + draft_manager = SimpleNamespace( + size=16, + dtype=torch.float32, + layer_num=3, + head_dim=2, + ) + manager.register_dflash_draft_mem_manager(draft_manager, global_kv_heads=2) + monkeypatch.setattr( + "lightllm.common.kv_cache_mem_manager.qwen3next_mem_manager.get_env_start_args", + lambda: SimpleNamespace(cpu_cache_token_page_size=page_size), + ) + + draft_bytes = 3 * page_size * 4 * 2 * torch.float32.itemsize + cpu_cache = torch.full((page_num, 1, 1, 1, target_bytes + draft_bytes), 0xA5, dtype=torch.uint8) + draft_cache = manager.get_dflash_draft_cpu_cache(cpu_cache) + + assert draft_cache.shape == (page_num, 3, page_size, 4, 2) + assert draft_cache.dtype == torch.float32 + draft_cache.zero_() + assert torch.all(cpu_cache.reshape(page_num, -1)[:, :target_bytes] == 0xA5) + assert torch.all(cpu_cache.reshape(page_num, -1)[:, target_bytes:] == 0) + + +def test_qwen35_dflash_does_not_reserve_duplicate_target_kv_layers(tmp_path, monkeypatch): + (tmp_path / "config.json").write_text( + '{"architectures": ["Qwen3_5DFlashModel"], "n_layer": 5}', + encoding="utf-8", + ) + args = SimpleNamespace( + mtp_mode="dflash", + mtp_step=15, + mtp_dynamic_verify=False, + mtp_draft_model_dir=[str(tmp_path)], + ) + monkeypatch.setattr(envs_utils, "get_env_start_args", lambda: args) + envs_utils.get_added_mtp_kv_layer_num.cache_clear() + + assert envs_utils.get_added_mtp_kv_layer_num() == 0 + envs_utils.get_added_mtp_kv_layer_num.cache_clear() + + +def test_qwen35_dflash_cpu_cache_meta_includes_draft_kv(monkeypatch): + args = SimpleNamespace( + enable_cpu_cache=True, + model_dir="unused-target", + mtp_mode="dflash", + mtp_step=15, + mtp_dynamic_verify=False, + cpu_cache_storage_size=1, + ) + monkeypatch.setattr(kv_cache_utils, "get_env_start_args", lambda: args) + monkeypatch.setattr(kv_cache_utils, "is_linear_att_mixed_model", lambda _path: True) + monkeypatch.setattr(kv_cache_utils, "get_llm_data_type", lambda: torch.bfloat16) + monkeypatch.setattr( + kv_cache_utils.LinearAttCacheConfig, + "load_from_args", + lambda: SimpleNamespace(get_cpu_cache_big_page_bytes=lambda: 1024), + ) + monkeypatch.setattr(kv_cache_utils, "_get_qwen35_dflash_cpu_cache_bytes", lambda _args: 384) + kv_cache_utils.calcu_cpu_cache_meta.cache_clear() + + meta = kv_cache_utils.calcu_cpu_cache_meta() + + assert meta.data_type == torch.uint8 + assert meta.calcu_one_page_size() == 1024 + 384 + kv_cache_utils.calcu_cpu_cache_meta.cache_clear() + + +def test_qwen35_dflash_cpu_cache_draft_size_uses_checkpoint_shape(monkeypatch): + from transformers.configuration_utils import PretrainedConfig + + monkeypatch.setattr( + PretrainedConfig, + "get_config_dict", + lambda *_args, **_kwargs: ({"architectures": ["Qwen3_5DFlashModel"]}, {}), + ) + monkeypatch.setattr(kv_cache_utils, "get_layer_num", lambda _path: 3) + monkeypatch.setattr(kv_cache_utils, "get_num_key_value_heads", lambda _path: 2) + monkeypatch.setattr(kv_cache_utils, "get_head_dim", lambda _path: 5) + monkeypatch.setattr(kv_cache_utils, "get_llm_data_type", lambda: torch.bfloat16) + args = SimpleNamespace(mtp_draft_model_dir=["unused-draft"], cpu_cache_token_page_size=4) + + draft_bytes = kv_cache_utils._get_qwen35_dflash_cpu_cache_bytes(args) + + assert draft_bytes == 4 * 3 * 2 * 2 * 5 * torch.bfloat16.itemsize From 464bc48ed0e49b5f0fac5552ad337f0e6667c1f2 Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Wed, 5 Aug 2026 06:23:43 +0000 Subject: [PATCH 003/103] remove dflash dynamic draft & add qwen35 dspark --- .../common/basemodel/attention/base_att.py | 5 + .../common/basemodel/attention/linear/gdn.py | 48 ++++- lightllm/common/basemodel/basemodel.py | 34 +++- lightllm/common/basemodel/batch_objs.py | 5 + lightllm/common/basemodel/cuda_graph.py | 25 ++- .../linear_att/mtp_state_params.py | 94 ++++++++++ .../operator/linear_att.py | 43 +---- .../qwen3next_mem_manager.py | 171 +----------------- .../linear_att_cache_manager/config_objs.py | 14 ++ .../{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json | 8 + .../{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json | 8 + .../{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json | 8 + .../{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json | 7 + ...ARLEN=true,REVERSE=false}_NVIDIA_H800.json | 38 ++++ ...=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json | 7 + ...2,a_dtype=torch.bfloat16}_NVIDIA_H800.json | 50 +++++ ...6,x_dtype=torch.bfloat16}_NVIDIA_H800.json | 50 +++++ ...M=6,dtype=torch.bfloat16}_NVIDIA_H800.json | 50 +++++ ...out_dtype=torch.bfloat16}_NVIDIA_H800.json | 74 ++++++++ lightllm/models/__init__.py | 1 + lightllm/models/qwen3_5_dflash/model.py | 123 +++++-------- lightllm/models/qwen3_5_dspark/__init__.py | 3 + lightllm/models/qwen3_5_dspark/model.py | 23 +++ .../layer_infer/post_layer_infer.py | 139 +++++++++++++- lightllm/server/pd_io_struct.py | 2 +- .../model_infer/mode_backend/base_backend.py | 54 ++++-- .../pd/decode_node_impl/decode_impl.py | 11 +- .../pd/prefill_node_impl/prefill_impl.py | 24 +-- .../router/model_infer/speculative/planner.py | 5 + .../speculative/proposers/dflash.py | 27 +-- .../speculative/proposers/dspark.py | 5 +- lightllm/utils/envs_utils.py | 6 +- lightllm/utils/kv_cache_utils.py | 41 +---- 33 files changed, 796 insertions(+), 407 deletions(-) create mode 100644 lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_gated_delta_rule_fwd_h/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_local_cumsum_scalar/{B=1,BT=64,H=12,IS_VARLEN=true,REVERSE=false}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_scaled_dot_kkt_fwd/{BT=64,H=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/fused_gdn_gating:v1/{NUM_HEADS=12,a_dtype=torch.bfloat16}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/gated_rmsnorm_forward:v1/{N=128,has_bias=false,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/mrope_triton_fused:v1/{HEAD_DIM=256,K_HEAD_NUM=1,Q_HEAD_NUM=6,dtype=torch.bfloat16}_NVIDIA_H800.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/silu_and_mul_fwd:v1/{N=4352,out_dtype=torch.bfloat16}_NVIDIA_H800.json create mode 100644 lightllm/models/qwen3_5_dspark/__init__.py create mode 100644 lightllm/models/qwen3_5_dspark/model.py diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index a1985a6fb9..bc1c39df20 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -104,6 +104,11 @@ class BaseDecodeAttState(ABC): backend: BaseAttBackend = None infer_state: "InferStateInfo" = None + def prepare_for_forward(self): + """Build derived state that must execute inside a captured forward.""" + + return + @abstractmethod def init_state(self): pass diff --git a/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index c4c3e15cb7..482d240283 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -2,7 +2,7 @@ import torch from typing import TYPE_CHECKING from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl -from lightllm.utils.envs_utils import get_env_start_args, get_llm_data_type +from lightllm.utils.envs_utils import enable_dynamic_mtp_verify, get_env_start_args, get_llm_data_type from lightllm.common.basemodel.triton_kernel.linear_att.causal_conv1d import causal_conv1d_fn from lightllm.common.basemodel.triton_kernel.linear_att.fused_gdn_gating import fused_gdn_gating from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import chunk_gated_delta_rule @@ -202,6 +202,47 @@ class LinearAttDecodeAttState(BaseDecodeAttState): b1_mtp_cu_q_seq_len: torch.Tensor = None b_num_accepted_tokens: torch.Tensor = None + def _uses_dynamic_mtp_layout(self) -> bool: + return ( + self.backend.mtp_step > 0 + and enable_dynamic_mtp_verify() + and not getattr(self.infer_state, "use_static_mtp_layout", False) + ) + + def prepare_for_forward(self): + """Build compact GDN row metadata as part of the captured forward. + + CUDA graph replay copies the primary ModelInput tensors into the graph + state. Deriving these tensors inside the graph makes every replay use + the current compact request layout instead of capture-time metadata. + """ + + if not self._uses_dynamic_mtp_layout(): + return + + from lightllm.common.basemodel.triton_kernel.linear_att.mtp_state_params import ( + build_dynamic_mtp_linear_att_state_params, + ) + + backend: LinearAttBackend = self.backend + batch_size = self.infer_state.batch_size + ( + self.b1_mtp_cu_q_seq_len, + self.b_conv_buffer_idx, + self.b_num_accepted_tokens, + ) = build_dynamic_mtp_linear_att_state_params( + b_req_idx=self.infer_state.b_req_idx, + b_mtp_index=self.infer_state.b_mtp_index, + req_to_mtp_state_index=self.infer_state.req_manager.req_to_mtp_state_index, + hold_req_id=self.infer_state.req_manager.HOLD_REQUEST_ID, + ) + self.b_ssm_buffer_idx = self.b_conv_buffer_idx.view(batch_size, 1) * (backend.mtp_step + 1) + torch.arange( + backend.mtp_step + 1, + device=self.infer_state.b_req_idx.device, + dtype=self.infer_state.b_req_idx.dtype, + ).view(1, backend.mtp_step + 1) + return + def init_state(self): backend: LinearAttBackend = self.backend mtp_step = backend.mtp_step @@ -216,6 +257,11 @@ def init_state(self): if mtp_step > 0: # mtp 模式下 batch_size = self.infer_state.batch_size + if self._uses_dynamic_mtp_layout(): + # This must run inside _token_forward so CUDA graph replay + # recomputes it from the current compact row layout. + return + att_batch_size = batch_size // (mtp_step + 1) assert batch_size % (mtp_step + 1) == 0 diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 3199082c66..7dcda28e41 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -84,7 +84,15 @@ def __init__(self, kvargs): assert not (self.is_token_healing and self.return_all_prompt_logics), "can not be true in same time" self.data_type = get_llm_data_type() mtp_step = get_env_start_args().mtp_step - self.graph_max_batch_size = kvargs.get("graph_max_batch_size", 16) + # A graph for more logical requests than the scheduler can ever run is + # unreachable. In MTP modes the value is expanded by ``mtp_step + 1`` + # below, so leaving the CLI default (often 256) uncapped can otherwise + # create multi-GiB full-vocabulary graph outputs for a server limited + # to only a few dozen concurrent requests. + self.graph_max_batch_size = min( + kvargs.get("graph_max_batch_size", 16), + self.max_req_num, + ) self.graph_max_batch_size = ( self.graph_max_batch_size // 2 if get_env_start_args().enable_decode_microbatch_overlap @@ -160,6 +168,15 @@ def set_spec_adapter(self, spec_adapter): self.spec_adapter = spec_adapter if self.graph is not None: self.graph.set_spec_adapter(spec_adapter, model=self) + # Models build their initial decode graphs before the speculative + # runtime exists. Attaching the adapter changes both graph batch + # sizes and cache keys, so eagerly capture the speculative graph + # variants on dummy state now. Lazy capture during the first live + # request can execute in-place KV/linear-state updates repeatedly. + if get_env_start_args().enable_decode_microbatch_overlap: + self.graph.warmup_overlap(self) + else: + self.graph.warmup(self) return def _init_config(self): @@ -558,6 +575,8 @@ def _create_unpad_decode_model_output( return model_output new_model_output = copy.copy(model_output) new_model_output.logits = new_model_output.logits[0:origin_batch_size] + if new_model_output.mtp_draft_token_ids is not None: + new_model_output.mtp_draft_token_ids = new_model_output.mtp_draft_token_ids[0:origin_batch_size] if new_model_output.mtp_draft_confidence_logits is not None: confidence_rows = new_model_output.mtp_draft_confidence_logits.shape[0] if confidence_rows == padded_batch_size: @@ -803,6 +822,13 @@ def prefill_func(input_tensors, infer_state): @final def _token_forward(self, infer_state: InferStateInfo): + # Some derived decode metadata depends on runtime tensor values and + # must therefore be computed inside CUDA graph capture/replay. The + # attention states cache it for reuse by every transformer layer. + infer_state.decode_att_state.prepare_for_forward() + if infer_state.decode_att_state1 is not None: + infer_state.decode_att_state1.prepare_for_forward() + input_ids = infer_state.input_ids cuda_input_ids = input_ids input_embs = self.pre_infer.token_forward(cuda_input_ids, infer_state, self.pre_post_weight) @@ -827,9 +853,15 @@ def _token_forward(self, infer_state: InferStateInfo): if pop_confidence_logits is not None: mtp_draft_confidence_logits = pop_confidence_logits() + mtp_draft_token_ids = None + pop_draft_token_ids = getattr(self.post_infer, "pop_mtp_draft_token_ids", None) + if pop_draft_token_ids is not None: + mtp_draft_token_ids = pop_draft_token_ids() + model_output = ModelOutput( logits=predict_logits.contiguous(), mtp_draft_confidence_logits=mtp_draft_confidence_logits, + mtp_draft_token_ids=mtp_draft_token_ids, ) if spec_context is not None: diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 7520c1aa45..66ce4d0196 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -128,6 +128,9 @@ class ModelOutput: # DSpark dynamic verify 使用的 raw confidence logits,由 draft model # post layer 产生,proposer 只负责 scatter 到 verify batch。 mtp_draft_confidence_logits: Optional[torch.Tensor] = None + # TP-sharded DSpark Markov head 已经完成全局 greedy 选择时,直接携带 + # draft token,避免为了下游 argmax 再 all-gather 完整词表 logits。 + mtp_draft_token_ids: Optional[torch.Tensor] = None # prompt_logics 用于在开启 return_all_prompt_logics 模式(如 enable_prompt_logprobs)时, # 保存整个 prefill 阶段每一个 token 位置对应的 logits(而非仅最后一个位置的 logits)。 @@ -139,3 +142,5 @@ def to_no_ref_tensor(self): self.logits = tensor_to_no_ref_tensor(self.logits) if self.mtp_draft_confidence_logits is not None: self.mtp_draft_confidence_logits = tensor_to_no_ref_tensor(self.mtp_draft_confidence_logits) + if self.mtp_draft_token_ids is not None: + self.mtp_draft_token_ids = tensor_to_no_ref_tensor(self.mtp_draft_token_ids) diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index c5f588be9f..46eb8fdf4d 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -1,3 +1,5 @@ +import gc + import torch import torch.distributed as dist import copy @@ -75,6 +77,18 @@ def __init__( return def set_spec_adapter(self, spec_adapter, model=None): + # Decode graphs captured before the speculative runtime is attached + # use different batch sizes and cache keys, so none of them can be + # replayed afterwards. Release their private graph pools before the + # speculative warmup; retaining both generations can consume several + # extra GiB and OOM otherwise valid mem_fraction configurations. + if self.graph: + torch.cuda.synchronize() + self.graph.clear() + self.mempool = None + gc.collect() + torch.cuda.empty_cache() + self.mempool = torch.cuda.graph_pool_handle() self.spec_adapter = spec_adapter self.model = model self._refresh_cuda_graph_batch_sizes() @@ -432,15 +446,16 @@ def warmup(self, model): del model_output if ( enable_dynamic_mtp_verify() - and self.args.mtp_mode == "eagle3" + and self.args.mtp_mode in {"eagle3", "dspark"} and self.spec_adapter is not None and not self.spec_adapter.is_draft_model(model) and batch_size % (self.args.mtp_step + 1) == 0 ): - # Dynamic Eagle3's profitable full-width state reuses the - # fixed K+1 FA3 layout. Capture that graph variant eagerly; - # otherwise its distinct spec key falls back to eager decode - # during the measurement and can be much slower than Static. + # Dynamic planners can switch back to the fixed K+1 target + # layout for a full-width iteration. Capture that graph + # variant on the dummy warmup state. Lazy capture on a live + # request would execute the in-place linear-attention state + # updates multiple times and corrupt subsequent decode state. model_input.use_static_mtp_layout = True model_output = model.forward(model_input) del model_output diff --git a/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py b/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py new file mode 100644 index 0000000000..a9704310ec --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py @@ -0,0 +1,94 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _build_dynamic_mtp_linear_att_state_params_kernel( + b_req_idx, + b_mtp_index, + req_to_mtp_state_index, + out_cu_q_seq_len, + out_conv_buffer_idx, + out_num_accepted_tokens, + batch_size, + hold_req_id, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.arange(0, BLOCK_SIZE) + token_mask = offsets < batch_size + cu_mask = offsets <= batch_size + + req_idx = tl.load(b_req_idx + offsets, mask=token_mask, other=hold_req_id) + mtp_index = tl.load(b_mtp_index + offsets, mask=token_mask, other=0) + + # CUDA-graph warmup uses HOLD_REQUEST_ID for every row but still supplies + # a real 0..K MTP layout. Runtime graph padding, by contrast, appends HOLD + # rows whose mtp_index is always zero. Preserve warmup work while excluding + # runtime padding from the final real sequence. + has_hold_draft_row = tl.sum( + tl.where(token_mask & (req_idx == hold_req_id) & (mtp_index > 0), 1, 0), + axis=0, + ) > 0 + valid_row = token_mask & ((req_idx != hold_req_id) | has_hold_draft_row) + actual_token_num = tl.sum(tl.where(valid_row, 1, 0), axis=0) + + # Keep tensor shapes dependent only on the target graph batch size. The + # unused tail represents zero-length sequences, so one captured graph can + # replay arbitrary compact per-request widths at the same token batch size. + tl.store(out_cu_q_seq_len + offsets, actual_token_num, mask=cu_mask) + tl.store(out_conv_buffer_idx + offsets, hold_req_id, mask=token_mask) + tl.store(out_num_accepted_tokens + offsets, 1, mask=token_mask) + tl.debug_barrier() + + is_start = valid_row & (mtp_index == 0) + sequence_index = tl.cumsum(tl.where(is_start, 1, 0), axis=0) - 1 + accepted_state_index = tl.load( + req_to_mtp_state_index + req_idx, + mask=is_start, + other=0, + ) + + tl.store(out_cu_q_seq_len + sequence_index, offsets, mask=is_start) + tl.store(out_conv_buffer_idx + sequence_index, req_idx, mask=is_start) + tl.store(out_num_accepted_tokens + sequence_index, accepted_state_index + 1, mask=is_start) + + +def build_dynamic_mtp_linear_att_state_params( + *, + b_req_idx: torch.Tensor, + b_mtp_index: torch.Tensor, + req_to_mtp_state_index: torch.Tensor, + hold_req_id: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Build graph-stable varlen GDN parameters for compact MTP verify rows. + + The compact input remains request-major and each request contributes a + prefix beginning at ``b_mtp_index == 0``. Outputs are padded to the token + batch size; trailing entries describe zero-length sequences. + """ + + assert b_req_idx.is_cuda and b_mtp_index.is_cuda and req_to_mtp_state_index.is_cuda + assert b_req_idx.ndim == 1 and b_req_idx.shape == b_mtp_index.shape + assert b_req_idx.dtype == torch.int32 and b_mtp_index.dtype == torch.int32 + batch_size = b_req_idx.shape[0] + assert batch_size > 0 + + b1_cu_q_seq_len = torch.empty((batch_size + 1,), dtype=torch.int32, device=b_req_idx.device) + b_conv_buffer_idx = torch.empty_like(b_req_idx) + b_num_accepted_tokens = torch.empty_like(b_req_idx) + + _build_dynamic_mtp_linear_att_state_params_kernel[(1,)]( + b_req_idx=b_req_idx, + b_mtp_index=b_mtp_index, + req_to_mtp_state_index=req_to_mtp_state_index, + out_cu_q_seq_len=b1_cu_q_seq_len, + out_conv_buffer_idx=b_conv_buffer_idx, + out_num_accepted_tokens=b_num_accepted_tokens, + batch_size=batch_size, + hold_req_id=int(hold_req_id), + BLOCK_SIZE=triton.next_power_of_2(batch_size + 1), + num_warps=8, + num_stages=1, + ) + return b1_cu_q_seq_len, b_conv_buffer_idx, b_num_accepted_tokens diff --git a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py index 6d5152e4a6..68b3b09a46 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py +++ b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py @@ -89,24 +89,6 @@ def load_cpu_cache_to_gpu( big_page_token_num=args.cpu_cache_token_page_size, linear_config=self.linear_config, ) - if mem_manager.has_separate_dflash_draft_kv: - from lightllm.common.basemodel.triton_kernel.kv_cache_offload import ( - load_cpu_kv_to_gpu, - ) - - draft_mem = mem_manager._pd_dflash_draft_mem_manager - draft_mem_indexes = mem_indexes.masked_fill(mem_indexes < 0, draft_mem.HOLD_TOKEN_MEMINDEX) - load_cpu_kv_to_gpu( - gpu_mem_indexes=draft_mem_indexes, - gpu_kv_cache=draft_mem.kv_buffer, - gpu_kv_cache_scale=None, - cpu_kv_cache=mem_manager.get_dflash_draft_cpu_cache(cpu_cache_client.cpu_kv_cache_tensor), - cpu_kv_cache_scale=None, - page_indexes=page_indexes, - tp_index=get_current_rank_in_dp(), - tp_world_size=get_dp_world_size(), - grid_num=16, - ) from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -205,31 +187,12 @@ def offload_gpu_kv_to_cpu_cache( big_page_token_num=args.cpu_cache_token_page_size, linear_config=self.linear_config, ) - if mem_manager.has_separate_dflash_draft_kv: - from lightllm.common.basemodel.triton_kernel.kv_cache_offload import ( - offload_gpu_kv_to_cpu, - ) - - draft_mem = mem_manager._pd_dflash_draft_mem_manager - draft_mem_indexes = mem_indexes.masked_fill(mem_indexes < 0, draft_mem.HOLD_TOKEN_MEMINDEX) - offload_gpu_kv_to_cpu( - token_indexes=draft_mem_indexes, - gpu_kv_cache=draft_mem.kv_buffer, - gpu_kv_cache_scale=None, - cpu_kv_cache=mem_manager.get_dflash_draft_cpu_cache(cpu_cache_client.cpu_kv_cache_tensor), - cpu_kv_cache_scale=None, - page_indexes=page_indexes, - page_readies=page_readies, - tp_index=get_current_rank_in_dp(), - tp_world_size=get_dp_world_size(), - grid_num=16, - ) - return def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): - # Qwen3Next 需要调整 layer_index - layer_index = layer_index // self.linear_config.full_attention_interval + # Qwen3Next packs main full-attention layers first, followed by + # speculative draft layers. + layer_index = self.linear_config.get_full_att_kv_layer_index(layer_index) from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager mem_manager: MemoryManager = self.mem_manager diff --git a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py index d37a226ea9..907cc494a6 100644 --- a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py @@ -2,7 +2,6 @@ import triton from lightllm.utils.log_utils import init_logger from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager -from lightllm.common.kv_trans_kernel.nixl_kv_trans import page_io from lightllm.utils.envs_utils import get_env_start_args from lightllm.common.linear_att_cache_manager import LinearAttCacheConfig, LinearAttCacheManager from .operator import LinearAttMemOperator @@ -26,31 +25,11 @@ def __init__( mem_fraction=0.9, ): self.linear_config = linear_config - self._pd_dflash_draft_mem_manager = None - self._pd_dflash_global_kv_heads = None super().__init__(size, dtype, num_kv_heads, head_dim, full_att_layer_num, always_copy, mem_fraction) - @property - def has_separate_dflash_draft_kv(self) -> bool: - return self._pd_dflash_draft_mem_manager is not None - - def register_dflash_draft_mem_manager(self, draft_mem_manager: MemoryManager, global_kv_heads: int): - """Attach the independent Qwen3.5 DFlash KV cache for PD transfer.""" - assert draft_mem_manager.size == self.size - assert draft_mem_manager.dtype == self.dtype - global_kv_heads = int(global_kv_heads) - assert global_kv_heads > 0 - self._pd_dflash_draft_mem_manager = draft_mem_manager - self._pd_dflash_global_kv_heads = global_kv_heads - return - def get_att_input_params(self, layer_index: int) -> Tuple[Any, Any]: - if layer_index >= self.linear_config.all_layer_num: - # MTP draft full-attn layers are packed after the main model layers. - layer_index -= self.linear_config.linear_layer_num - else: - layer_index = layer_index // self.linear_config.full_attention_interval + layer_index = self.linear_config.get_full_att_kv_layer_index(layer_index) return super().get_att_input_params(layer_index) def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): @@ -101,146 +80,10 @@ def write_to_shm(self, req_manager): self.linear_att_big_page_buffers = big_page_buffers def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: - if not self.has_separate_dflash_draft_kv: - kv_move_buffer = super().alloc_paged_kv_move_buffer(page_num, page_size) - else: - # PD pages are opaque contiguous bytes to the transporter. Allocate - # one registered page large enough for either target or draft KV; - # the page-kind hooks below expose the corresponding shaped view. - target_elements = ( - self.layer_num * 2 * self.linear_config.full_att_all_num_kv_heads * self.head_dim - ) - draft_mem = self._pd_dflash_draft_mem_manager - draft_elements = ( - draft_mem.layer_num * 2 * self._pd_dflash_global_kv_heads * draft_mem.head_dim - ) - linear_state_nbytes = Qwen3NextLinearAttPageHelper(self).state_nbytes - linear_state_elements_per_token = ( - linear_state_nbytes + page_size * self.dtype.itemsize - 1 - ) // (page_size * self.dtype.itemsize) - elements_per_token = max( - target_elements, - draft_elements, - linear_state_elements_per_token, - ) - self.kv_move_buffer = torch.empty( - (page_num, page_size, 1, 1, elements_per_token), - dtype=self.dtype, - device="cuda", - ) - self._buffer_mem_indexes_tensors = [ - torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) - for _ in range(page_num) - ] - kv_move_buffer = self.kv_move_buffer + kv_move_buffer = super().alloc_paged_kv_move_buffer(page_num, page_size) Qwen3NextLinearAttPageHelper(self).assert_page_size() return kv_move_buffer - def _get_pd_kv_page(self, page_index: int, page_kind: str) -> torch.Tensor: - page = self.kv_move_buffer[page_index] - if not self.has_separate_dflash_draft_kv: - assert page_kind == "kv" - return page - - flat_page = page.reshape(-1) - if page_kind == "kv": - layer_num = self.layer_num - global_kv_heads = self.linear_config.full_att_all_num_kv_heads - head_dim = self.head_dim - elif page_kind == "draft_kv": - draft_mem = self._pd_dflash_draft_mem_manager - layer_num = draft_mem.layer_num - global_kv_heads = self._pd_dflash_global_kv_heads - head_dim = draft_mem.head_dim - else: - raise ValueError(f"unknown KV page kind {page_kind}") - - element_num = layer_num * 2 * global_kv_heads * head_dim - total_element_num = page.shape[0] * element_num - assert total_element_num <= flat_page.numel() - return flat_page[:total_element_num].view(page.shape[0], layer_num, 2 * global_kv_heads, head_dim) - - def _get_pd_mem_manager(self, page_kind: str) -> MemoryManager: - if page_kind == "kv": - return self - if page_kind == "draft_kv" and self.has_separate_dflash_draft_kv: - return self._pd_dflash_draft_mem_manager - raise ValueError(f"unknown KV page kind {page_kind}") - - def get_dflash_draft_cpu_cache(self, cpu_cache_tensor: torch.Tensor) -> torch.Tensor: - """Return the draft-KV portion appended to a mixed-model CPU cache page.""" - - assert self.has_separate_dflash_draft_kv - assert cpu_cache_tensor.dtype == torch.uint8 - draft_mem = self._pd_dflash_draft_mem_manager - page_num = cpu_cache_tensor.shape[0] - page_size = get_env_start_args().cpu_cache_token_page_size - draft_shape = ( - page_num, - draft_mem.layer_num, - page_size, - 2 * self._pd_dflash_global_kv_heads, - draft_mem.head_dim, - ) - draft_nbytes = 1 - for dim in draft_shape[1:]: - draft_nbytes *= dim - draft_nbytes *= self.dtype.itemsize - - # Existing Qwen3Next full-attention/linear-state data occupies the - # aligned prefix. Qwen3.5 DFlash KV is stored in the remaining bytes so - # disk offload and shared-memory page lifecycle stay unchanged. - draft_offset = self.linear_config.get_cpu_cache_big_page_bytes() - flat_cache = cpu_cache_tensor.reshape(page_num, -1) - assert draft_offset + draft_nbytes <= flat_cache.shape[1] - draft_bytes = flat_cache[:, draft_offset : draft_offset + draft_nbytes] - return draft_bytes.view(dtype=self.dtype).view(draft_shape) - - def _write_kv_mem_to_page( - self, mem_indexes, page_index: int, dp_index: int, mem_managers, dp_world_size: int, page_kind: str - ): - page = self._get_pd_kv_page(page_index, page_kind) - pin_mem_indexes = self._buffer_mem_indexes_tensors[page_index][0 : len(mem_indexes)] - pin_mem_indexes.numpy()[:] = mem_indexes - mem_indexes_gpu = pin_mem_indexes.cuda(non_blocking=True) - dp_mems = mem_managers[(dp_index * dp_world_size) : ((dp_index + 1) * dp_world_size)] - assert len(dp_mems) == dp_world_size - source_mems = [mem._get_pd_mem_manager(page_kind) for mem in dp_mems] - repeat_count = dp_world_size * source_mems[0].kv_buffer.shape[2] // page.shape[2] - assert repeat_count > 0 - for tp_index, mem in enumerate(source_mems): - if tp_index % repeat_count == 0: - page_io( - mem_indexes=mem_indexes_gpu, - page_tensor=page, - kv_buffer=mem.kv_buffer, - tp_index=tp_index, - tp_world_size=dp_world_size, - mode="write", - ) - return - - def _read_kv_page_to_mem( - self, mem_indexes, page_index: int, dp_index: int, mem_managers, dp_world_size: int, page_kind: str - ): - page = self._get_pd_kv_page(page_index, page_kind) - pin_mem_indexes = self._buffer_mem_indexes_tensors[page_index][0 : len(mem_indexes)] - pin_mem_indexes.numpy()[:] = mem_indexes - mem_indexes_gpu = pin_mem_indexes.cuda(non_blocking=True) - dp_mems = mem_managers[(dp_index * dp_world_size) : ((dp_index + 1) * dp_world_size)] - assert len(dp_mems) == dp_world_size - target_mems = [mem._get_pd_mem_manager(page_kind) for mem in dp_mems] - for tp_index, mem in enumerate(target_mems): - page_io( - mem_indexes=mem_indexes_gpu, - page_tensor=page, - kv_buffer=mem.kv_buffer, - tp_index=tp_index, - tp_world_size=dp_world_size, - mode="read", - ) - return - def write_mem_to_page_kv_move_buffer( self, mem_indexes, @@ -251,14 +94,15 @@ def write_mem_to_page_kv_move_buffer( page_kind: str = "kv", req_idx: int = None, ): - if page_kind in ("kv", "draft_kv"): - return self._write_kv_mem_to_page( + if page_kind == "kv": + return super().write_mem_to_page_kv_move_buffer( mem_indexes=mem_indexes, page_index=page_index, dp_index=dp_index, mem_managers=mem_managers, dp_world_size=dp_world_size, page_kind=page_kind, + req_idx=req_idx, ) assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" assert req_idx is not None @@ -277,14 +121,15 @@ def read_page_kv_move_buffer_to_mem( page_kind: str = "kv", req_idx: int = None, ): - if page_kind in ("kv", "draft_kv"): - return self._read_kv_page_to_mem( + if page_kind == "kv": + return super().read_page_kv_move_buffer_to_mem( mem_indexes=mem_indexes, page_index=page_index, dp_index=dp_index, mem_managers=mem_managers, dp_world_size=dp_world_size, page_kind=page_kind, + req_idx=req_idx, ) assert page_kind == "linear_att_state", f"unknown page_kind={page_kind}" assert req_idx is not None diff --git a/lightllm/common/linear_att_cache_manager/config_objs.py b/lightllm/common/linear_att_cache_manager/config_objs.py index b63cd6b0e7..f588ec7d5c 100644 --- a/lightllm/common/linear_att_cache_manager/config_objs.py +++ b/lightllm/common/linear_att_cache_manager/config_objs.py @@ -50,6 +50,20 @@ def get_main_model_full_att_layer_num(self): def get_full_att_kv_layer_num_with_draft_model(self): return self.get_main_model_full_att_layer_num() + self.draft_full_att_kv_layer_num + def get_full_att_kv_layer_index(self, layer_index: int) -> int: + """Map a global target/draft layer index to the packed full-attention cache.""" + + layer_index = int(layer_index) + if layer_index >= self.all_layer_num: + kv_layer_index = layer_index - self.linear_layer_num + else: + kv_layer_index = layer_index // self.full_attention_interval + assert 0 <= kv_layer_index < self.get_full_att_kv_layer_num_with_draft_model(), ( + f"layer {layer_index} maps outside the packed full-attention KV cache: " + f"slot={kv_layer_index}, slots={self.get_full_att_kv_layer_num_with_draft_model()}" + ) + return kv_layer_index + def get_conv_state_shape(self): # Base committed sliding-window state, without speculative MTP tail. return (self.get_conv_dim(), self.conv_kernel_size - 1) diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json new file mode 100644 index 0000000000..f55b637832 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json @@ -0,0 +1,8 @@ +{ + "4": { + "BK": 128, + "BV": 128, + "num_stages": 4, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json new file mode 100644 index 0000000000..cc5c68eb79 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json @@ -0,0 +1,8 @@ +{ + "4": { + "BK": 128, + "BV": 64, + "num_stages": 2, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json new file mode 100644 index 0000000000..7421097fa4 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json @@ -0,0 +1,8 @@ +{ + "4": { + "BK": 64, + "BV": 128, + "num_stages": 3, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_gated_delta_rule_fwd_h/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_gated_delta_rule_fwd_h/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json new file mode 100644 index 0000000000..d831f32c4a --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_gated_delta_rule_fwd_h/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json @@ -0,0 +1,7 @@ +{ + "4": { + "BV": 32, + "num_stages": 4, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_local_cumsum_scalar/{B=1,BT=64,H=12,IS_VARLEN=true,REVERSE=false}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_local_cumsum_scalar/{B=1,BT=64,H=12,IS_VARLEN=true,REVERSE=false}_NVIDIA_H800.json new file mode 100644 index 0000000000..14509fffea --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_local_cumsum_scalar/{B=1,BT=64,H=12,IS_VARLEN=true,REVERSE=false}_NVIDIA_H800.json @@ -0,0 +1,38 @@ +{ + "1": { + "num_warps": 1 + }, + "1024": { + "num_warps": 2 + }, + "128": { + "num_warps": 2 + }, + "16": { + "num_warps": 8 + }, + "2048": { + "num_warps": 1 + }, + "256": { + "num_warps": 2 + }, + "32": { + "num_warps": 1 + }, + "4": { + "num_warps": 2 + }, + "4096": { + "num_warps": 2 + }, + "64": { + "num_warps": 1 + }, + "8": { + "num_warps": 1 + }, + "8192": { + "num_warps": 2 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_scaled_dot_kkt_fwd/{BT=64,H=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_scaled_dot_kkt_fwd/{BT=64,H=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json new file mode 100644 index 0000000000..a97cabf8b2 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_scaled_dot_kkt_fwd/{BT=64,H=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json @@ -0,0 +1,7 @@ +{ + "4": { + "BK": 64, + "num_stages": 3, + "num_warps": 2 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/fused_gdn_gating:v1/{NUM_HEADS=12,a_dtype=torch.bfloat16}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/fused_gdn_gating:v1/{NUM_HEADS=12,a_dtype=torch.bfloat16}_NVIDIA_H800.json new file mode 100644 index 0000000000..5284108618 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/fused_gdn_gating:v1/{NUM_HEADS=12,a_dtype=torch.bfloat16}_NVIDIA_H800.json @@ -0,0 +1,50 @@ +{ + "1": { + "BLK_HEADS": 64, + "num_warps": 2 + }, + "1024": { + "BLK_HEADS": 16, + "num_warps": 1 + }, + "128": { + "BLK_HEADS": 64, + "num_warps": 4 + }, + "16": { + "BLK_HEADS": 16, + "num_warps": 1 + }, + "2048": { + "BLK_HEADS": 64, + "num_warps": 2 + }, + "256": { + "BLK_HEADS": 16, + "num_warps": 1 + }, + "32": { + "BLK_HEADS": 8, + "num_warps": 2 + }, + "4": { + "BLK_HEADS": 4, + "num_warps": 1 + }, + "4096": { + "BLK_HEADS": 64, + "num_warps": 2 + }, + "64": { + "BLK_HEADS": 4, + "num_warps": 4 + }, + "8": { + "BLK_HEADS": 16, + "num_warps": 1 + }, + "8192": { + "BLK_HEADS": 16, + "num_warps": 1 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/gated_rmsnorm_forward:v1/{N=128,has_bias=false,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/gated_rmsnorm_forward:v1/{N=128,has_bias=false,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H800.json new file mode 100644 index 0000000000..233215c4f2 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/gated_rmsnorm_forward:v1/{N=128,has_bias=false,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H800.json @@ -0,0 +1,50 @@ +{ + "12": { + "BLOCK_N": 128, + "num_warps": 1 + }, + "12288": { + "BLOCK_N": 128, + "num_warps": 1 + }, + "1536": { + "BLOCK_N": 512, + "num_warps": 1 + }, + "192": { + "BLOCK_N": 512, + "num_warps": 1 + }, + "24576": { + "BLOCK_N": 512, + "num_warps": 1 + }, + "3072": { + "BLOCK_N": 128, + "num_warps": 1 + }, + "384": { + "BLOCK_N": 64, + "num_warps": 2 + }, + "48": { + "BLOCK_N": 256, + "num_warps": 2 + }, + "49152": { + "BLOCK_N": 128, + "num_warps": 1 + }, + "768": { + "BLOCK_N": 64, + "num_warps": 2 + }, + "96": { + "BLOCK_N": 256, + "num_warps": 2 + }, + "98304": { + "BLOCK_N": 128, + "num_warps": 1 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/mrope_triton_fused:v1/{HEAD_DIM=256,K_HEAD_NUM=1,Q_HEAD_NUM=6,dtype=torch.bfloat16}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/mrope_triton_fused:v1/{HEAD_DIM=256,K_HEAD_NUM=1,Q_HEAD_NUM=6,dtype=torch.bfloat16}_NVIDIA_H800.json new file mode 100644 index 0000000000..3e5f0d7165 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/mrope_triton_fused:v1/{HEAD_DIM=256,K_HEAD_NUM=1,Q_HEAD_NUM=6,dtype=torch.bfloat16}_NVIDIA_H800.json @@ -0,0 +1,50 @@ +{ + "1": { + "num_stages": 2, + "num_warps": 4 + }, + "1024": { + "num_stages": 4, + "num_warps": 4 + }, + "128": { + "num_stages": 2, + "num_warps": 1 + }, + "16": { + "num_stages": 4, + "num_warps": 4 + }, + "2048": { + "num_stages": 4, + "num_warps": 4 + }, + "256": { + "num_stages": 2, + "num_warps": 1 + }, + "32": { + "num_stages": 1, + "num_warps": 2 + }, + "4": { + "num_stages": 4, + "num_warps": 4 + }, + "4096": { + "num_stages": 2, + "num_warps": 2 + }, + "64": { + "num_stages": 2, + "num_warps": 2 + }, + "8": { + "num_stages": 4, + "num_warps": 4 + }, + "8192": { + "num_stages": 2, + "num_warps": 2 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/silu_and_mul_fwd:v1/{N=4352,out_dtype=torch.bfloat16}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/silu_and_mul_fwd:v1/{N=4352,out_dtype=torch.bfloat16}_NVIDIA_H800.json new file mode 100644 index 0000000000..6fc416c0ec --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/silu_and_mul_fwd:v1/{N=4352,out_dtype=torch.bfloat16}_NVIDIA_H800.json @@ -0,0 +1,74 @@ +{ + "1": { + "BLOCK_M": 128, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "1024": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "128": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "16": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2048": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "256": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "32": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 8 + }, + "4": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 4 + }, + "4096": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "64": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "8": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "8192": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + } +} \ No newline at end of file diff --git a/lightllm/models/__init__.py b/lightllm/models/__init__.py index 7a9aa3cd8f..bd8db7f8dd 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -65,6 +65,7 @@ "Qwen3_5TpPartModel": ("lightllm.models.qwen3_5.model", "Qwen3_5TpPartModel"), "Qwen3_5MOETpPartModel": ("lightllm.models.qwen3_5_moe.model", "Qwen3_5MOETpPartModel"), "Qwen3_5DFlashModel": ("lightllm.models.qwen3_5_dflash.model", "Qwen3_5DFlashModel"), + "Qwen3_5DSparkModel": ("lightllm.models.qwen3_5_dspark.model", "Qwen3_5DSparkModel"), } _MODEL_TYPE_REGISTRY_MODULES = { diff --git a/lightllm/models/qwen3_5_dflash/model.py b/lightllm/models/qwen3_5_dflash/model.py index 571fc1f1e7..e27e23dae5 100644 --- a/lightllm/models/qwen3_5_dflash/model.py +++ b/lightllm/models/qwen3_5_dflash/model.py @@ -1,5 +1,3 @@ -from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager -from lightllm.common.kv_cache_mem_manager.operator import NormalMemOperator from lightllm.common.speculative.config import normalize_speculative_draft_config from lightllm.distributed.communication_op import dist_group_manager from lightllm.models.llama.model import LlamaTpPartModel @@ -9,43 +7,6 @@ ) -class Qwen35DFlashMemOperator(NormalMemOperator): - """Map global draft layer indexes onto the draft-owned KV cache. - - Qwen3.5 target models use a Qwen3Next mixed linear/full-attention memory - manager, which is not compatible with ordinary DFlash KV writes. The - Qwen3 DFlash layer infer code is reused here and still passes global layer - indexes, so this operator converts those indexes back to the local - 0..draft_layer_num-1 range used by the independent draft MemoryManager. - """ - - def _local_layer_index(self, layer_index: int) -> int: - return int(layer_index) - int(self.mem_manager.layer_index_offset) - - def copy_kv_to_mem_manager(self, layer_index: int, mem_index, kv): - return super().copy_kv_to_mem_manager(self._local_layer_index(layer_index), mem_index, kv) - - -class Qwen35DFlashMemoryManager(MemoryManager): - """Ordinary KV cache for the Qwen3.5 DFlash draft model. - - Unlike Qwen3 DFlash, this draft cannot share the target model memory - manager because the Qwen3.5 target cache contains Qwen3Next linear-attention - state. The layer_index_offset records where the draft layers sit in the - global model layer list, while the underlying cache stores only draft-local - layer rows. - """ - - operator_class = Qwen35DFlashMemOperator - - def __init__(self, *args, layer_index_offset: int = 0, **kwargs): - self.layer_index_offset = int(layer_index_offset) - super().__init__(*args, **kwargs) - - def get_att_input_params(self, layer_index: int): - return super().get_att_input_params(int(layer_index) - self.layer_index_offset) - - class Qwen3_5DFlashModel(Qwen3DFlashModel): """Qwen3.5 DFlash draft model. @@ -56,21 +17,7 @@ class Qwen3_5DFlashModel(Qwen3DFlashModel): """ pre_and_post_weight_class = Qwen35DFlashPreAndPostLayerWeight - - def __init__(self, kvargs: dict): - super().__init__(kvargs) - # The request manager is shared with the target, while this draft owns - # a separate KV cache. TpPartBaseModel wires every request manager to - # the model's memory manager during construction, which would make - # request allocation bypass the target/radix-cache allocator. Restore - # the shared request manager's allocator after draft initialization. - self._restore_main_mem_manager() - return - - def _restore_main_mem_manager(self): - assert self.req_manager is self.main_model.req_manager - self.req_manager.mem_manager = self.main_model.mem_manager - return + share_target_embedding_and_lm_head = True def _init_config(self): super()._init_config() @@ -84,8 +31,13 @@ def _normalize_dflash_config(self): if isinstance(rope_parameters, dict): if "rope_theta" in rope_parameters and "rope_theta" not in self.config: self.config["rope_theta"] = rope_parameters["rope_theta"] - if "partial_rotary_factor" in rope_parameters and "partial_rotary_factor" not in self.config: - self.config["partial_rotary_factor"] = rope_parameters["partial_rotary_factor"] + if ( + "partial_rotary_factor" in rope_parameters + and "partial_rotary_factor" not in self.config + ): + self.config["partial_rotary_factor"] = rope_parameters[ + "partial_rotary_factor" + ] if "rope_scaling" not in self.config: self.config["rope_scaling"] = rope_parameters return @@ -101,35 +53,42 @@ def _init_custom(self): return def _init_mem_manager(self): - head_dim = self.config.get("head_dim") - if head_dim is None: - head_dim = self.config["hidden_size"] // self.config["num_attention_heads"] - num_kv_heads = max(1, self.config["num_key_value_heads"] // self.tp_world_size_) - self.draft_layer_start = len(self.main_model.layers_infer) - self.draft_layer_start += sum( - len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models - ) - self.mem_manager = Qwen35DFlashMemoryManager( - size=self.main_model.mem_manager.size, - dtype=self.data_type, - head_num=num_kv_heads, - head_dim=head_dim, - layer_num=self.config["n_layer"], - mem_fraction=self.mem_fraction, - layer_index_offset=self.draft_layer_start, + target_mem_manager = self.main_model.mem_manager + target_linear_config = getattr(target_mem_manager, "linear_config", None) + assert ( + target_linear_config is not None + ), "Qwen3_5DFlashModel requires a Qwen3Next target" + + draft_head_dim = self.config.get("head_dim") + if draft_head_dim is None: + draft_head_dim = ( + self.config["hidden_size"] // self.config["num_attention_heads"] + ) + draft_kv_heads = int(self.config["num_key_value_heads"]) + target_kv_heads = int(target_linear_config.full_att_all_num_kv_heads) + draft_layer_num = int(self.config["n_layer"]) + reserved_draft_layer_num = int(target_linear_config.draft_full_att_kv_layer_num) + assert ( + int(draft_head_dim) == int(target_mem_manager.head_dim) + and draft_kv_heads == target_kv_heads + ), ( + "Qwen3.5 block draft currently requires draft and target full-attention KV shapes to match: " + f"draft=({draft_kv_heads}, {draft_head_dim}), " + f"target=({target_kv_heads}, {target_mem_manager.head_dim})" ) - register_draft_manager = getattr( - self.main_model.mem_manager, "register_dflash_draft_mem_manager", None + assert reserved_draft_layer_num >= draft_layer_num, ( + "Qwen3Next target KV cache did not reserve enough draft layers: " + f"required={draft_layer_num}, reserved={reserved_draft_layer_num}" ) - assert callable(register_draft_manager), "Qwen3_5DFlashModel requires a Qwen3Next target" - register_draft_manager( - self.mem_manager, - global_kv_heads=int(self.config["num_key_value_heads"]), - ) - return + return super()._init_mem_manager() def _init_weights(self, start_layer_index=None): super()._init_weights(start_layer_index=start_layer_index) - self.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ - self.pre_post_weight.lm_head_weight_ = self.main_model.pre_post_weight.lm_head_weight_ + if self.share_target_embedding_and_lm_head: + self.pre_post_weight.wte_weight_ = ( + self.main_model.pre_post_weight.wte_weight_ + ) + self.pre_post_weight.lm_head_weight_ = ( + self.main_model.pre_post_weight.lm_head_weight_ + ) return diff --git a/lightllm/models/qwen3_5_dspark/__init__.py b/lightllm/models/qwen3_5_dspark/__init__.py new file mode 100644 index 0000000000..ee28c7b584 --- /dev/null +++ b/lightllm/models/qwen3_5_dspark/__init__.py @@ -0,0 +1,3 @@ +from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel + +__all__ = ["Qwen3_5DSparkModel"] diff --git a/lightllm/models/qwen3_5_dspark/model.py b/lightllm/models/qwen3_5_dspark/model.py new file mode 100644 index 0000000000..102b8a0d76 --- /dev/null +++ b/lightllm/models/qwen3_5_dspark/model.py @@ -0,0 +1,23 @@ +from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel +from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import ( + Qwen3DSparkPostLayerInfer, +) +from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import ( + Qwen3DSparkPreAndPostLayerWeight, +) + + +class Qwen3_5DSparkModel(Qwen3_5DFlashModel): + """DSpark draft model paired with a Qwen3.5 hybrid-attention target. + + DeepSpec exports these checkpoints as ``Qwen3DSparkModel`` because the + draft backbone itself is a stack of Qwen3 full-attention layers. The + target pairing still matters at serving time: Qwen3.5 uses draft-owned + rotary caches and a target-owned compatible KV cache, while proposal logits + use DSpark's Markov/confidence heads and the ordinary zero-based block + layout. + """ + + pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight + post_layer_infer_class = Qwen3DSparkPostLayerInfer + share_target_embedding_and_lm_head = False diff --git a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py index e5efa6768f..4a9490359d 100644 --- a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py +++ b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py @@ -2,7 +2,7 @@ import torch import torch.nn.functional as F -from lightllm.distributed.communication_op import all_gather +from lightllm.distributed.communication_op import all_gather, all_gather_into_tensor from lightllm.models.llama.infer_struct import LlamaInferStateInfo from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import ( @@ -26,12 +26,18 @@ def __init__(self, network_config): self.enable_confidence_head_ = bool(network_config.get("enable_confidence_head", False)) self.confidence_head_with_markov_ = bool(network_config.get("confidence_head_with_markov", False)) self.mtp_draft_confidence_logits = None + self.mtp_draft_token_ids = None def pop_mtp_draft_confidence_logits(self): logits = self.mtp_draft_confidence_logits self.mtp_draft_confidence_logits = None return logits + def pop_mtp_draft_token_ids(self): + token_ids = self.mtp_draft_token_ids + self.mtp_draft_token_ids = None + return token_ids + def has_markov_head(self) -> bool: return self.markov_rank_ > 0 @@ -180,6 +186,90 @@ def _token_forward_with_hidden( gather_data = None return logits, head_hidden + def _token_forward_with_local_logits_and_hidden( + self, + input_embdings: torch.Tensor, + infer_state: LlamaInferStateInfo, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ): + """Project the LM head but keep its vocabulary shard local. + + Vanilla Markov decoding only needs the global maximum token at each + block position. Keeping logits sharded avoids a full-vocabulary TP + all-gather and prevents every rank from repeating the same Markov + projection over the complete vocabulary. + """ + + last_input, token_num = self._slice_get_last_input(input_embdings, infer_state) + head_hidden = last_input + normed_input = self._norm(last_input, infer_state, layer_weight) + lm_head_input = normed_input.permute(1, 0).reshape(-1, token_num) + local_logits = layer_weight.lm_head_weight_(input=lm_head_input, alloc_func=self.alloc_tensor) + return local_logits, head_hidden + + def _sample_tp_sharded_vanilla_markov( + self, + local_logits: torch.Tensor, + *, + infer_state: LlamaInferStateInfo, + anchor_token_ids: torch.Tensor, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ) -> torch.Tensor: + """Run exact greedy Markov decoding on TP vocabulary shards.""" + + assert self.markov_head_type_ == "vanilla" + assert layer_weight.markov_w2_weight_ is not None + token_num = local_logits.shape[1] + assert token_num % self.block_size_ == 0 + num_reqs = token_num // self.block_size_ + + vocab_size = int(layer_weight.lm_head_weight_.vocab_size) + split_indexes = np.linspace(0, vocab_size, self.tp_world_size_ + 1, dtype=np.int64) + local_start = int(split_indexes[self.tp_rank_]) + local_end = int(split_indexes[self.tp_rank_ + 1]) + assert local_logits.shape[0] == local_end - local_start, ( + f"local LM head rows must match TP vocabulary shard [{local_start}, {local_end}), " + f"got {local_logits.shape[0]}" + ) + local_markov_w2 = layer_weight.markov_w2_weight_.weight[local_start:local_end, :] + + prev_token_ids = anchor_token_ids.long() + sampled_tokens = [] + req_rows = torch.arange(num_reqs, dtype=torch.long, device=local_logits.device) + for step_idx in range(self.block_size_): + prev_embeddings = self._markov_prev_embeddings(prev_token_ids, layer_weight) + local_markov_bias = F.linear(prev_embeddings.to(dtype=local_markov_w2.dtype), local_markov_w2) + local_base_logits = local_logits[:, step_idx::self.block_size_].permute(1, 0).float() + local_scores = local_base_logits + local_markov_bias + local_max_values, local_max_indexes = torch.max(local_scores, dim=-1) + local_token_ids = local_max_indexes + local_start + + if self.tp_world_size_ == 1: + next_token_ids = local_token_ids + else: + local_winners = torch.stack( + [local_max_values, local_token_ids.to(dtype=torch.float32)], + dim=-1, + ).contiguous() + gathered_winners = self.alloc_tensor( + (self.tp_world_size_ * num_reqs, 2), + dtype=torch.float32, + ) + all_gather_into_tensor( + gathered_winners, + local_winners, + group=infer_state.dist_group, + async_op=False, + ) + gathered_winners = gathered_winners.view(self.tp_world_size_, num_reqs, 2) + winning_ranks = torch.argmax(gathered_winners[:, :, 0], dim=0) + next_token_ids = gathered_winners[winning_ranks, req_rows, 1].long() + + sampled_tokens.append(next_token_ids) + prev_token_ids = next_token_ids + + return torch.stack(sampled_tokens, dim=1) + def token_forward( self, input_embdings: torch.Tensor, @@ -187,6 +277,7 @@ def token_forward( layer_weight: Qwen3DSparkPreAndPostLayerWeight, ): self.mtp_draft_confidence_logits = None + self.mtp_draft_token_ids = None if self._is_commit_prefill(infer_state): return super().token_forward( input_embdings=input_embdings, @@ -194,13 +285,55 @@ def token_forward( layer_weight=layer_weight, ) + if infer_state.is_prefill: + logits, _ = self._token_forward_with_hidden( + input_embdings=input_embdings, + infer_state=infer_state, + layer_weight=layer_weight, + ) + return logits + + use_tp_sharded_markov = ( + self.tp_world_size_ > 1 + and self.has_markov_head() + and self.markov_head_type_ == "vanilla" + ) + if use_tp_sharded_markov: + local_logits, head_hidden = self._token_forward_with_local_logits_and_hidden( + input_embdings=input_embdings, + infer_state=infer_state, + layer_weight=layer_weight, + ) + token_num = local_logits.shape[1] + assert token_num % self.block_size_ == 0 + num_reqs = token_num // self.block_size_ + block_hidden = head_hidden.reshape(num_reqs, self.block_size_, -1) + anchor_token_ids = infer_state.input_ids.reshape(num_reqs, self.block_size_)[:, 0] + sampled_tokens = self._sample_tp_sharded_vanilla_markov( + local_logits, + infer_state=infer_state, + anchor_token_ids=anchor_token_ids, + layer_weight=layer_weight, + ) + self.mtp_draft_token_ids = sampled_tokens.reshape(-1) + self.mtp_draft_confidence_logits = self.predict_confidence_logits( + block_hidden, + anchor_token_ids=anchor_token_ids, + sampled_tokens=sampled_tokens, + layer_weight=layer_weight, + ) + # The proposer consumes mtp_draft_token_ids directly. Keep the + # leading row dimension for generic graph padding/unpadding while + # avoiding an otherwise unused [rows, vocab] tensor. A single + # placeholder column is required because CUDA graph's no-ref + # tensor wrapper cannot represent a zero-byte allocation. + return local_logits.new_empty((token_num, 1)) + logits, head_hidden = self._token_forward_with_hidden( input_embdings=input_embdings, infer_state=infer_state, layer_weight=layer_weight, ) - if infer_state.is_prefill: - return logits assert ( logits.shape[0] % self.block_size_ == 0 diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 524d1a97b4..d0f32419d6 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -184,7 +184,7 @@ def __post_init__(self): error_info = "start_kv_index must >=0 and end_kv_index > start_kv_index" logger.error(error_info) raise ValueError(error_info) - if self.page_kind in ("kv", "draft_kv"): + if self.page_kind == "kv": assert len(self.mem_indexes) == (self.end_kv_index - self.start_kv_index) elif self.page_kind == "linear_att_state": assert self.start_kv_index == self.end_kv_index 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 961044b757..926baabff7 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -252,6 +252,12 @@ def init_model(self, kvargs): [rank for rank in range(self.global_world_size)], backend="nccl" ) + if self.args.run_mode in ["prefill", "decode"] or self.args.enable_dp_prompt_cache_fetch: + # The target manager already includes the speculative full-attention + # layer slots, so it can be shared before draft model initialization. + self.model.mem_manager.write_to_shm(req_manager=self.model.req_manager) + dist.barrier(group=self.node_nccl_group) + # 同一 DP 组内只需主 rank 初始化真实的 capture buffer 并执行后续相关操作; # 非主 rank 不需要分配 buffer,避免重复占用内存。 if self.is_master_in_dp: @@ -267,6 +273,9 @@ def init_model(self, kvargs): self.init_custom() + if self.args.enable_dp_prompt_cache_fetch: + self.init_dp_kv_shared() + self.shm_reqs_io_buffer = ShmObjsIOBuffer() # 只会在 pd pd 模式下才会使用,用于上传分块传输任务是否成功。 self.shm_pd_trans_io_buffer = ShmObjsIOBuffer(tail_str="pd") @@ -279,16 +288,6 @@ def init_model(self, kvargs): self.spec_adapter = build_spec_runtime(self) self._attach_spec_adapter() - if self.args.run_mode in ["prefill", "decode"] or self.args.enable_dp_prompt_cache_fetch: - # Draft models must be initialized before this snapshot. Qwen3.5 - # DFlash attaches its independent draft KV manager to the target - # manager so PD transfer workers can deserialize both buffers. - self.model.mem_manager.write_to_shm(req_manager=self.model.req_manager) - dist.barrier(group=self.node_nccl_group) - - if self.args.enable_dp_prompt_cache_fetch: - self.init_dp_kv_shared() - if self.args.enable_cpu_cache: self.multi_level_cache_module = MultiLevelKvCacheModule(self) @@ -438,9 +437,14 @@ def init_mtp_draft_model(self, main_kvargs: dict): self.draft_models.append(Qwen3DFlashModel(mtp_model_kvargs)) elif spec_config.is_dspark and is_qwen3_dspark_draft_config(mtp_model_cfg): - from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel + if self.is_linear_att_mixed_model: + from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel - self.draft_models.append(Qwen3DSparkModel(mtp_model_kvargs)) + self.draft_models.append(Qwen3_5DSparkModel(mtp_model_kvargs)) + else: + from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel + + self.draft_models.append(Qwen3DSparkModel(mtp_model_kvargs)) elif (spec_config.is_dflash or spec_config.is_dspark) and is_gemma4_dspark_draft_config(mtp_model_cfg): raise NotImplementedError("Gemma4 DSpark draft checkpoints are not wired to LightLLM serving yet.") elif (spec_config.is_dflash or spec_config.is_dspark) and is_dspark_draft_config(mtp_model_cfg): @@ -477,9 +481,9 @@ def _normalize_block_mtp_step_from_config(self, mtp_model_cfg: dict) -> None: return def _validate_linear_att_spec_support(self) -> None: - """Restrict the new DFlash combination without narrowing existing LightSpec modes.""" + """Validate block-draft checkpoint families against hybrid targets.""" - if not self.spec_config.enabled or not self.spec_config.is_dflash: + if not self.spec_config.enabled or not self.spec_config.uses_block_draft_model: return mtp_draft_model_dirs = self.args.mtp_draft_model_dir @@ -488,11 +492,18 @@ def _validate_linear_att_spec_support(self) -> None: assert mtp_draft_model_dirs is not None and len(mtp_draft_model_dirs) == 1 mtp_model_cfg, _ = PretrainedConfig.get_config_dict(mtp_draft_model_dirs[0]) is_qwen35_dflash = is_qwen3_5_dflash_draft_config(mtp_model_cfg) + is_qwen_dspark = is_qwen3_dspark_draft_config(mtp_model_cfg) if self.is_linear_att_mixed_model: - assert is_qwen35_dflash, ( - "linear-attention mixed targets require a Qwen3_5DFlashModel draft checkpoint, " - f"got architectures={mtp_model_cfg.get('architectures')}" - ) + if self.spec_config.is_dflash: + assert is_qwen35_dflash, ( + "linear-attention mixed targets require a Qwen3_5DFlashModel checkpoint in DFlash mode, " + f"got architectures={mtp_model_cfg.get('architectures')}" + ) + else: + assert is_qwen_dspark, ( + "linear-attention mixed targets require a Qwen3DSparkModel checkpoint in DSpark mode, " + f"got architectures={mtp_model_cfg.get('architectures')}" + ) else: assert not is_qwen35_dflash, "Qwen3_5DFlashModel requires a Qwen3Next target" return @@ -1053,8 +1064,11 @@ def _update_mtp_verify_token_num( return def _gen_argmax_token_ids(self, model_output: ModelOutput): - logits = model_output.logits - draft_next_token_ids_gpu = torch.argmax(logits, dim=-1) + if model_output.mtp_draft_token_ids is not None: + draft_next_token_ids_gpu = model_output.mtp_draft_token_ids + else: + logits = model_output.logits + draft_next_token_ids_gpu = torch.argmax(logits, dim=-1) # 如果draft和target的词表不同,需要把draft token映射回主模型词表。 if self.spec_config.needs_draft_vocab_mapping: diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 1a20fc4f82..f9dc6ee60f 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -144,15 +144,6 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): kv_end_index=end_index, group=group, ) - if getattr(self.model.mem_manager, "has_separate_dflash_draft_kv", False): - self._create_pd_trans_task( - req_obj=req_obj, - mem_indexes=page_mem_indexes.tolist(), - kv_start_index=start_index, - kv_end_index=end_index, - group=group, - page_kind="draft_kv", - ) # update req_obj.pd_trans_kv_start_index += cur_page_size @@ -203,7 +194,7 @@ def _create_pd_trans_task( # only self.is_master_in_dp will be used. self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size - if page_kind in ("kv", "draft_kv"): + if page_kind == "kv": req_idx = None elif page_kind == "linear_att_state": req_idx = req_obj.req_idx diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index 4fbb1eaf89..2a501f509b 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -63,25 +63,13 @@ def _prefill_chuncked_handle_func( cur_page_size = min(page_size, req_obj.cur_kv_len - req_obj.pd_trans_kv_start_index) # 生成页面传输任务, 放入kv move manager 的处理队列中 if cur_page_size == page_size or prefill_finished: - start_index = req_obj.pd_trans_kv_start_index - end_index = start_index + cur_page_size - trans_task_list.append( - self._create_pd_trans_task( - req_obj=req_obj, - kv_start_index=start_index, - kv_end_index=end_index, - ) + trans_task = self._create_pd_trans_task( + req_obj=req_obj, + kv_start_index=req_obj.pd_trans_kv_start_index, + kv_end_index=req_obj.pd_trans_kv_start_index + cur_page_size, ) - if getattr(self.model.mem_manager, "has_separate_dflash_draft_kv", False): - trans_task_list.append( - self._create_pd_trans_task( - req_obj=req_obj, - kv_start_index=start_index, - kv_end_index=end_index, - page_kind="draft_kv", - ) - ) req_obj.pd_trans_kv_start_index += cur_page_size + trans_task_list.append(trans_task) else: break @@ -118,7 +106,7 @@ def _create_pd_trans_task( self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size pd_decode_node_info = req_obj.sampling_param.pd_decode_node - if page_kind in ("kv", "draft_kv"): + if page_kind == "kv": mem_indexes = ( self.model.req_manager.req_to_token_indexs[req_obj.req_idx, kv_start_index:kv_end_index] .detach() diff --git a/lightllm/server/router/model_infer/speculative/planner.py b/lightllm/server/router/model_infer/speculative/planner.py index 21fd1bdbfa..66969663e2 100644 --- a/lightllm/server/router/model_infer/speculative/planner.py +++ b/lightllm/server/router/model_infer/speculative/planner.py @@ -1346,6 +1346,11 @@ def update_predicted_schedule_probs(self, *, schedule_probs, req_num: int) -> No def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: assert req_num * (self.mtp_step + 1) == original_batch_size + # ``DynamicMTPPlanner.plan`` rewrites a full-width confidence plan to + # ``observe`` for that iteration. DSpark may select a narrower plan on + # the next iteration, so do not carry the previous iteration's mode + # into the new capacity decision. + self._selection_mode = "confidence" pre_draft_step = self.pre_draft_step self.pre_draft_step = self.mtp_step if req_num == 0: diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index 4936f89d65..0393d28344 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -80,10 +80,7 @@ def propose_next( draft_probs=None, ) - self.extend_draft_kv_cache( - main_model_input=main_model_input, - accepted_index=verify_result.accepted_index, - ) + self.extend_draft_kv_cache(main_model_input=main_model_input) # DFlash only drafts from the accepted tail row of each request. Unlike # MTP, one anchor row expands to a whole non-causal block. @@ -125,15 +122,11 @@ def select_draft_token_ids( assert 0 <= output_start <= output_end <= block_token_ids.shape[1] return block_token_ids[:, output_start:output_end] - def extend_draft_kv_cache(self, *, main_model_input: ModelInput, accepted_index: torch.Tensor) -> None: + def extend_draft_kv_cache(self, *, main_model_input: ModelInput) -> None: target_hidden = self.runtime.get_hidden() draft_model = self.backend.draft_models[0] - accepted_rows = torch.nonzero(accepted_index.to(torch.bool), as_tuple=False).flatten().to(torch.long) - if accepted_rows.numel() == 0: - return - target_hidden = target_hidden.index_select(0, accepted_rows) - batch_size = int(target_hidden.shape[0]) + assert batch_size == main_model_input.b_req_idx.shape[0] draft_kv_input = copy.copy(main_model_input) draft_kv_input.batch_size = batch_size @@ -141,7 +134,7 @@ def extend_draft_kv_cache(self, *, main_model_input: ModelInput, accepted_index: draft_kv_input.multimodal_params = [{"images": [], "audios": []} for _ in range(batch_size)] # This hidden-commit prefill path does not consume token ids, but # InferState uses input_ids.shape[0] to build position ids. Keep it - # aligned with the accepted hidden rows. + # aligned with the fixed-shape target verify batch. draft_kv_input.input_ids = torch.empty( (batch_size,), dtype=torch.int64, @@ -150,12 +143,12 @@ def extend_draft_kv_cache(self, *, main_model_input: ModelInput, accepted_index: draft_kv_input.max_q_seq_len = 1 draft_kv_input.prefix_total_token_num = 0 draft_kv_input.is_prefill = True - # Each accepted MTP row writes one target-hidden KV slot for the same request. - draft_kv_input.b_req_idx = main_model_input.b_req_idx.index_select(0, accepted_rows).contiguous() - draft_kv_input.b_mtp_index = main_model_input.b_mtp_index.index_select(0, accepted_rows).contiguous() - draft_kv_input.b_seq_len = main_model_input.b_seq_len.index_select(0, accepted_rows).contiguous() - draft_kv_input.mem_indexes = main_model_input.mem_indexes.index_select(0, accepted_rows).contiguous() - draft_kv_input.b_ready_cache_len = draft_kv_input.b_seq_len - 1 + # Match Eagle3's fixed verify commit: write every speculative row and + # let the accepted-tail sequence length select the valid prefix. Rejected + # suffix slots are ignored and released by the normal verify free path. + # Keeping the batch shape static avoids torch.nonzero's implicit D2H + # synchronization and lets the host enqueue the draft work immediately. + draft_kv_input.b_ready_cache_len = main_model_input.b_seq_len - 1 draft_kv_input.b_prefill_start_loc = torch.arange( batch_size, dtype=torch.int32, diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/speculative/proposers/dspark.py index aca931f759..1b32a162c0 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dspark.py @@ -55,10 +55,7 @@ def propose_next( schedule_probs=schedule_probs, ) - self.extend_draft_kv_cache( - main_model_input=main_model_input, - accepted_index=verify_result.accepted_index, - ) + self.extend_draft_kv_cache(main_model_input=main_model_input) selected_rows = self.select_accepted_tail_rows( b_req_mtp_start_loc=b_req_mtp_start_loc, accept_len=verify_result.accept_len, diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 119e21cb7c..b16e1612ac 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -5,7 +5,7 @@ from easydict import EasyDict from functools import lru_cache from lightllm.utils.log_utils import init_logger -from lightllm.common.speculative import SpeculativeConfig, is_qwen3_5_dflash_draft_config +from lightllm.common.speculative import SpeculativeConfig logger = init_logger(__name__) @@ -279,10 +279,6 @@ def get_added_mtp_kv_layer_num() -> int: return spec_config.draft_model_count with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: draft_config = json.load(json_file) - if spec_config.is_dflash and is_qwen3_5_dflash_draft_config(draft_config): - # Qwen3.5 DFlash owns a separate KV manager; do not reserve duplicate - # draft layers in the mixed target manager. - return 0 return int(draft_config.get("num_hidden_layers", draft_config.get("n_layer", spec_config.draft_model_count))) return spec_config.draft_model_count diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index 43864aa56b..6856382996 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -16,7 +16,7 @@ get_added_mtp_kv_layer_num, ) from lightllm.utils.log_utils import init_logger -from lightllm.common.speculative import SpeculativeConfig, is_qwen3_5_dflash_draft_config +from lightllm.common.speculative import SpeculativeConfig from lightllm.utils.config_utils import get_num_key_value_heads, get_head_dim, get_layer_num, is_linear_att_mixed_model from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class from lightllm.common.kv_cache_mem_manager import ( @@ -59,35 +59,6 @@ def compute_token_list_hash(tokens: List[int], cpu_cache_token_page_size: int) - return chunks_hash_value -def _get_qwen35_dflash_cpu_cache_bytes(args) -> int: - """Size of the global draft KV appended to each Qwen3Next CPU cache page.""" - - draft_model_dirs = args.mtp_draft_model_dir - if isinstance(draft_model_dirs, str): - draft_model_dirs = [draft_model_dirs] - assert draft_model_dirs and len(draft_model_dirs) == 1 - - from transformers.configuration_utils import PretrainedConfig - - draft_config, _ = PretrainedConfig.get_config_dict(draft_model_dirs[0]) - assert is_qwen3_5_dflash_draft_config(draft_config), ( - "linear-attention MTP CPU cache only supports Qwen3_5DFlashModel, " - f"got architectures={draft_config.get('architectures')}" - ) - page_size = args.cpu_cache_token_page_size - draft_bytes = ( - page_size - * get_layer_num(draft_model_dirs[0]) - * 2 - * get_num_key_value_heads(draft_model_dirs[0]) - * get_head_dim(draft_model_dirs[0]) - * get_llm_data_type().itemsize - ) - # The following typed draft view requires its byte offset and extent to be - # naturally aligned. Existing mixed-model pages use the same alignment. - return triton.cdiv(draft_bytes, 16) * 16 - - @lru_cache(maxsize=None) def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": args = get_env_start_args() @@ -149,13 +120,9 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": raise Exception(f"not support mem manager: {mem_manager_class} for cpu kv cache") spec_config = SpeculativeConfig.from_args(args) - if spec_config.enabled: - if mem_manager_class is Qwen3NextMemManager: - if spec_config.is_dflash: - cpu_cache_meta.head_dim += _get_qwen35_dflash_cpu_cache_bytes(args) - else: - # TODO 可能会存在不同mtp模式的精度问题 - cpu_cache_meta.layer_num += get_added_mtp_kv_layer_num() + if spec_config.enabled and mem_manager_class is not Qwen3NextMemManager: + # TODO 可能会存在不同mtp模式的精度问题 + cpu_cache_meta.layer_num += get_added_mtp_kv_layer_num() cpu_cache_page_num = int( (args.cpu_cache_storage_size * 1024 * 1024 * 1024) / (cpu_cache_meta.calcu_one_page_size()) From 8a75a1f4078650308e258417bc31c70cff1c535b Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Mon, 10 Aug 2026 07:56:22 +0000 Subject: [PATCH 004/103] WIP: simplify LightSpec implementation (not ready for review) --- .../common/basemodel/attention/base_att.py | 19 +- lightllm/common/basemodel/attention/fa3/fp.py | 92 +-- .../common/basemodel/attention/fa3/fp8.py | 12 +- .../common/basemodel/attention/fa3/mla.py | 92 +-- .../common/basemodel/attention/linear/gdn.py | 127 +-- .../common/basemodel/attention/triton/fp.py | 42 +- lightllm/common/basemodel/basemodel.py | 354 +++----- lightllm/common/basemodel/batch_objs.py | 50 +- lightllm/common/basemodel/cuda_graph.py | 350 ++------ lightllm/common/basemodel/hidden_collector.py | 130 +++ lightllm/common/basemodel/infer_struct.py | 30 +- .../common/basemodel/prefill_cuda_graph.py | 8 +- ...mic_mtp_utils.py => dynamic_spec_utils.py} | 180 +++- .../basemodel/triton_kernel/fa3_utils.py | 89 +- ...p_state_params.py => spec_state_params.py} | 51 +- .../triton_kernel/linear_att_copy.py | 14 +- .../basemodel/triton_kernel/mtp_utils.py | 220 ++--- .../basemodel/triton_kernel/norm/qk_norm.py | 4 +- .../operator/linear_att.py | 3 +- .../linear_att_cache_manager/config_objs.py | 4 +- lightllm/common/req_manager.py | 36 +- lightllm/common/speculative/__init__.py | 27 - lightllm/common/speculative/config.py | 310 ------- ....bfloat16,q_head_dim=128}_NVIDIA_H800.json | 326 -------- ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 254 ------ ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 254 ------ ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 26 - ...h.float16,q_head_dim=128}_NVIDIA_H200.json | 26 - ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 26 - ...h.float16,q_head_dim=128}_NVIDIA_H200.json | 26 - ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 26 - ...h.float16,q_head_dim=128}_NVIDIA_H200.json | 26 - ....bfloat16,q_head_dim=128}_NVIDIA_H200.json | 26 - ...h.float16,q_head_dim=128}_NVIDIA_H200.json | 26 - .../{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json | 8 - .../{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json | 8 - .../{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json | 8 - .../{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json | 7 - ...ARLEN=true,REVERSE=false}_NVIDIA_H800.json | 38 - ...=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json | 7 - ...2,a_dtype=torch.bfloat16}_NVIDIA_H800.json | 50 -- ...6,x_dtype=torch.bfloat16}_NVIDIA_H800.json | 50 -- ...M=6,dtype=torch.bfloat16}_NVIDIA_H800.json | 50 -- ...out_dtype=torch.bfloat16}_NVIDIA_H800.json | 74 -- lightllm/models/__init__.py | 273 +++--- lightllm/models/deepseek_mtp/model.py | 7 +- lightllm/models/glm4_moe_lite_mtp/model.py | 7 +- lightllm/models/mistral_mtp/model.py | 7 +- .../pre_and_post_layer_weight.py | 8 - lightllm/models/qwen3_5_dflash/model.py | 52 +- lightllm/models/qwen3_5_dspark/model.py | 3 +- lightllm/models/qwen3_dflash/infer_struct.py | 30 +- .../qwen3_dflash/layer_infer/__init__.py | 9 - .../layer_infer/post_layer_infer.py | 19 +- .../layer_infer/pre_layer_infer.py | 61 +- .../layer_infer/transformer_layer_infer.py | 129 +-- .../pre_and_post_layer_weight.py | 1 - .../layer_weights/transformer_layer_weight.py | 22 +- lightllm/models/qwen3_dflash/model.py | 66 +- .../qwen3_dspark/layer_infer/__init__.py | 4 - .../layer_infer/post_layer_infer.py | 48 +- .../pre_and_post_layer_weight.py | 1 - lightllm/models/qwen3_dspark/model.py | 43 +- lightllm/models/qwen3_dspark/model_output.py | 20 + .../qwen3_eagle/layer_infer/__init__.py | 7 - .../layer_infer/pre_layer_infer.py | 32 +- .../layer_infer/transformer_layer_infer.py | 31 +- .../pre_and_post_layer_weight.py | 7 +- .../layer_weights/transformer_layer_weight.py | 64 +- lightllm/models/qwen3_eagle/model.py | 37 +- lightllm/models/qwen3_moe_mtp/model.py | 7 +- lightllm/server/api_cli.py | 37 +- lightllm/server/api_start.py | 74 +- lightllm/server/core/objs/req.py | 1 + lightllm/server/core/objs/start_args_type.py | 3 - lightllm/server/httpserver/manager.py | 42 +- .../httpserver_for_pd_master/manager.py | 9 +- .../server/router/model_infer/infer_batch.py | 71 +- .../model_infer/mode_backend/__init__.py | 61 +- .../model_infer/mode_backend/base_backend.py | 337 +++----- .../mode_backend/chunked_prefill/impl.py | 62 +- .../mode_backend/diverse_backend/impl.py | 13 +- .../mode_backend/dp_backend/impl.py | 393 +++++---- .../generic_padded_pre_process.py | 44 +- .../mode_backend/generic_post_process.py | 149 +--- .../mode_backend/generic_pre_process.py | 56 +- .../mode_backend/update_mem_index.py | 50 -- .../model_infer/speculative/__init__.py | 42 +- .../router/model_infer/speculative/engine.py | 531 ++++++++++++ .../router/model_infer/speculative/planner.py | 697 +++------------- .../speculative/proposers/__init__.py | 58 +- .../model_infer/speculative/proposers/base.py | 86 +- .../speculative/proposers/dflash.py | 67 +- .../speculative/proposers/dspark.py | 58 +- .../speculative/proposers/eagle3.py | 67 +- .../speculative/proposers/eagle_mtp.py | 405 +++++++-- .../speculative/proposers/vanilla_mtp.py | 47 +- .../router/model_infer/speculative/runner.py | 88 +- .../router/model_infer/speculative/runtime.py | 775 ------------------ .../router/model_infer/speculative/state.py | 130 --- .../model_infer/speculative/verifier.py | 29 +- lightllm/utils/envs_utils.py | 40 +- lightllm/utils/kv_cache_utils.py | 13 +- lightllm/utils/sgl_utils.py | 15 +- test/speculative/test_qwen35_dflash_state.py | 378 --------- .../basemodel/test_cuda_graph_layout.py | 71 ++ .../common/basemodel/test_hidden_collector.py | 149 ++++ .../common/basemodel/test_model_output.py | 40 + .../test_int8kv_flash_decoding_diverse.py | 12 +- ...tp_utils.py => test_dynamic_spec_utils.py} | 120 +-- .../basemodel/triton_kernel/test_fa3_utils.py | 16 +- .../basemodel/triton_kernel/test_mtp_utils.py | 61 +- unit_tests/common/speculative/test_config.py | 40 - .../models/test_qwen3_dspark_model_output.py | 59 ++ .../mode_backend/test_dp_spec_engine.py | 107 +++ .../mode_backend/test_generic_post_process.py | 6 +- .../mode_backend/test_generic_pre_process.py | 62 ++ .../speculative/test_eagle_overlap.py | 88 ++ .../model_infer/speculative/test_planner.py | 20 +- .../server/test_api_start_spec_config.py | 42 - unit_tests/utils/test_speculative_utils.py | 208 +++++ 121 files changed, 3876 insertions(+), 6734 deletions(-) create mode 100644 lightllm/common/basemodel/hidden_collector.py rename lightllm/common/basemodel/triton_kernel/{dynamic_mtp_utils.py => dynamic_spec_utils.py} (51%) rename lightllm/common/basemodel/triton_kernel/linear_att/{mtp_state_params.py => spec_state_params.py} (62%) delete mode 100644 lightllm/common/speculative/__init__.py delete mode 100644 lightllm/common/speculative/config.py delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H800/_fwd_kernel_mtp_diverse_stage1_single_token:v1/{block_seq=256,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=8,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_gated_delta_rule_fwd_h/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_local_cumsum_scalar/{B=1,BT=64,H=12,IS_VARLEN=true,REVERSE=false}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_scaled_dot_kkt_fwd/{BT=64,H=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/fused_gdn_gating:v1/{NUM_HEADS=12,a_dtype=torch.bfloat16}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/gated_rmsnorm_forward:v1/{N=128,has_bias=false,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/mrope_triton_fused:v1/{HEAD_DIM=256,K_HEAD_NUM=1,Q_HEAD_NUM=6,dtype=torch.bfloat16}_NVIDIA_H800.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/silu_and_mul_fwd:v1/{N=4352,out_dtype=torch.bfloat16}_NVIDIA_H800.json create mode 100644 lightllm/models/qwen3_dspark/model_output.py delete mode 100644 lightllm/server/router/model_infer/mode_backend/update_mem_index.py create mode 100644 lightllm/server/router/model_infer/speculative/engine.py delete mode 100644 lightllm/server/router/model_infer/speculative/runtime.py delete mode 100644 lightllm/server/router/model_infer/speculative/state.py delete mode 100644 test/speculative/test_qwen35_dflash_state.py create mode 100644 unit_tests/common/basemodel/test_cuda_graph_layout.py create mode 100644 unit_tests/common/basemodel/test_hidden_collector.py create mode 100644 unit_tests/common/basemodel/test_model_output.py rename unit_tests/common/basemodel/triton_kernel/{test_dynamic_mtp_utils.py => test_dynamic_spec_utils.py} (63%) delete mode 100644 unit_tests/common/speculative/test_config.py create mode 100644 unit_tests/models/test_qwen3_dspark_model_output.py create mode 100644 unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py create mode 100644 unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py create mode 100644 unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py delete mode 100644 unit_tests/server/test_api_start_spec_config.py create mode 100644 unit_tests/utils/test_speculative_utils.py diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index bc1c39df20..063cd3ccf1 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -1,7 +1,9 @@ import torch from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import TYPE_CHECKING, Tuple, Union, Dict +from typing import Optional, TYPE_CHECKING, Tuple, Union, Dict + +from lightllm.utils.envs_utils import enable_dynamic_spec, get_env_start_args if TYPE_CHECKING: from lightllm.common.basemodel.basemodel import TpPartBaseModel @@ -22,7 +24,7 @@ def __new__(cls, *args, **kwargs): 和缓存布局,不能只按 backend class 共享实例。 """ model = kwargs.get("model", args[0] if args else None) - instance_key = (cls, id(model)) + instance_key = (cls, model) if instance_key not in cls._instances: instance = super().__new__(cls) cls._instances[instance_key] = instance @@ -37,6 +39,14 @@ def create_att_prefill_state(self) -> "BasePrefillAttState": def create_att_decode_state(self) -> "BaseDecodeAttState": raise NotImplementedError("not impl") + def uses_dynamic_spec_verify_layout(self, infer_state: "InferStateInfo") -> bool: + if infer_state.draft_step == 0 or not enable_dynamic_spec(): + return False + + # Target verification may compact each request to a different row count. + # Block draft forwards still use their checkpoint-defined fixed layout. + return get_env_start_args().mtp_mode not in ("dspark", "dflash") or not self.model.is_mtp_draft_model + def _find_layer_index( self, k: torch.Tensor, v: torch.Tensor, att_state: Union["BasePrefillAttState", "BaseDecodeAttState"] ) -> int: @@ -104,11 +114,6 @@ class BaseDecodeAttState(ABC): backend: BaseAttBackend = None infer_state: "InferStateInfo" = None - def prepare_for_forward(self): - """Build derived state that must execute inside a captured forward.""" - - return - @abstractmethod def init_state(self): pass diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index 76cdd83044..89fdc61389 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -1,11 +1,11 @@ import dataclasses import torch from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl +from typing import Optional, TYPE_CHECKING from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.sgl_utils import flash_attn_with_kvcache, flash_attn_with_kvcache_autotune -from lightllm.utils.envs_utils import enable_dynamic_mtp_verify, get_env_start_args from lightllm.common.basemodel.triton_kernel.fa3_utils import ( - build_dynamic_mtp_fa3_decode_params, + build_dynamic_spec_fa3_decode_params, page_table_copy, ) from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor @@ -104,7 +104,7 @@ def _nomarl_prefill_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.infer_state.max_q_seq_len, softmax_scale=sm_scale, - causal=getattr(self.infer_state, "prefill_causal", True), + causal=self.infer_state.prefill_causal, window_size=window_size, softcap=0.0, k_descale=k_descale, @@ -126,53 +126,48 @@ class Fa3DecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: Fa3AttBackend = self.backend - args_mtp_step = getattr(self.infer_state, "decode_mtp_step", None) - if args_mtp_step is None: - args_mtp_step = get_env_start_args().mtp_step - if self.infer_state.disable_mtp_decode_att: - args_mtp_step = 0 - is_block_draft_decode = getattr(self.infer_state, "is_draft_model", False) - is_dynamic_mtp = ( - args_mtp_step > 0 - and enable_dynamic_mtp_verify() - and not is_block_draft_decode - and not getattr(self.infer_state, "use_static_mtp_layout", False) - ) + draft_step = self.infer_state.draft_step + decode_rows_per_request = draft_step + 1 + uses_dynamic_spec_verify_layout = self.backend.uses_dynamic_spec_verify_layout(self.infer_state) + + if draft_step > 0 and not uses_dynamic_spec_verify_layout: + assert self.infer_state.batch_size % decode_rows_per_request == 0, ( + "FA3 fixed-layout decode requires batch_size to be divisible by draft_step + 1, " + f"got batch_size={self.infer_state.batch_size}, draft_step={draft_step}." + ) - if is_dynamic_mtp: - att_batch_size = self.infer_state.batch_size - (b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len,) = build_dynamic_mtp_fa3_decode_params( + # 修正 mtp 在 fa3 下的输入。 + if uses_dynamic_spec_verify_layout: + (b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len,) = build_dynamic_spec_fa3_decode_params( b_req_idx=self.infer_state.b_req_idx, b_seq_len=self.infer_state.b_seq_len, b_mark_shared_group=self.infer_state.b_mark_shared_group, - att_batch_size=att_batch_size, + att_batch_size=self.infer_state.batch_size, hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, ) - b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) - self.cu_seqlens_q = b1_cu_q_seq_len.int() - self.cu_seqlens_k = b1_cu_kv_seq_len.int() - elif args_mtp_step > 0: - # 修正 mtp 在 fa3 下的输入。 - mtp_size = args_mtp_step + 1 + elif draft_step > 0: b_q_seq_len = torch.full( - (self.infer_state.b_seq_len.shape[0] // mtp_size,), - fill_value=mtp_size, + (self.infer_state.b_seq_len.shape[0] // decode_rows_per_request,), + fill_value=decode_rows_per_request, dtype=torch.int32, device=self.infer_state.b_seq_len.device, ) - b_kv_seq_len = self.infer_state.b_seq_len[mtp_size - 1 :: mtp_size] + b_kv_seq_len = self.infer_state.b_seq_len[draft_step::decode_rows_per_request] + b_att_req_idx = self.infer_state.b_req_idx[draft_step::decode_rows_per_request] + self.b_att_seq_len = b_kv_seq_len.contiguous() + else: + b_att_req_idx = self.infer_state.b_req_idx + self.b_att_seq_len = self.infer_state.b_seq_len + + if draft_step > 0: b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) self.cu_seqlens_q = b1_cu_q_seq_len.int() self.cu_seqlens_k = b1_cu_kv_seq_len.int() - att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) else: self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() - att_batch_size = self.infer_state.batch_size - - if not is_dynamic_mtp: - assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 + att_batch_size = b_att_req_idx.shape[0] model = self.backend.model # 可以使用 cuda graph的时候从 buffer中申请 if ( @@ -190,29 +185,12 @@ def init_state(self): device=self.infer_state.input_ids.device, ) - if is_dynamic_mtp: - page_table_copy( - page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], - req_to_token_indexs=model.req_manager.req_to_token_indexs, - b_req_idx=b_att_req_idx, - ) - self.decode_max_q_seq_len = args_mtp_step + 1 - elif args_mtp_step > 0: - page_table_copy( - page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], - req_to_token_indexs=model.req_manager.req_to_token_indexs, - b_req_idx=self.infer_state.b_req_idx[args_mtp_step :: (args_mtp_step + 1)], - ) - self.b_att_seq_len = self.infer_state.b_seq_len[args_mtp_step :: (args_mtp_step + 1)].contiguous() - self.decode_max_q_seq_len = args_mtp_step + 1 - else: - page_table_copy( - page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], - req_to_token_indexs=model.req_manager.req_to_token_indexs, - b_req_idx=self.infer_state.b_req_idx, - ) - self.b_att_seq_len = self.infer_state.b_seq_len - self.decode_max_q_seq_len = 1 + page_table_copy( + page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], + req_to_token_indexs=model.req_manager.req_to_token_indexs, + b_req_idx=b_att_req_idx, + ) + self.decode_max_q_seq_len = decode_rows_per_request return def copy_for_decode_cuda_graph(self, new_state: "Fa3DecodeAttState"): @@ -266,7 +244,7 @@ def _normal_decode_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, softmax_scale=sm_scale, - causal=getattr(self.infer_state, "decode_causal", True), + causal=self.infer_state.decode_causal, window_size=window_size, softcap=0.0, k_descale=k_descale, diff --git a/lightllm/common/basemodel/attention/fa3/fp8.py b/lightllm/common/basemodel/attention/fa3/fp8.py index dd0222ff5b..dea9ed1e0d 100644 --- a/lightllm/common/basemodel/attention/fa3/fp8.py +++ b/lightllm/common/basemodel/attention/fa3/fp8.py @@ -1,9 +1,11 @@ import dataclasses import torch from ..base_att import AttControl +from typing import Optional, TYPE_CHECKING from lightllm.utils.sgl_utils import flash_attn_with_kvcache from lightllm.common.basemodel.triton_kernel.quantization.q_per_head_fp8_quant import q_per_head_fp8_quant from lightllm.utils.vllm_utils import HAS_VLLM, vllm_ops +from typing import Union from .fp import Fa3AttBackend, Fa3PrefillAttState, Fa3DecodeAttState if HAS_VLLM: @@ -96,7 +98,7 @@ def _fp8_prefill_att( cu_seqlens_q=self.cu_seqlens_q, cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.infer_state.max_q_seq_len, - causal=getattr(self.infer_state, "prefill_causal", True), + causal=self.infer_state.prefill_causal, window_size=(-1, -1), softcap=0.0, q_descale=q_scale, @@ -116,7 +118,7 @@ def init_state(self): super().init_state() self.backend: Fp8Fa3AttBackend = self.backend - batch_size = self.page_table.shape[0] + att_batch_size = self.b_att_seq_len.shape[0] mem_manager = self.backend.model.mem_manager offline_scales: torch.Tensor = mem_manager.scales @@ -124,10 +126,10 @@ def init_state(self): # 为了减少推理计算量,在推理外部初始化k_descale和v_descale self.k_descale = ( - offline_scales[:, :head_num].view(-1, 1, head_num).expand(offline_scales.shape[0], batch_size, head_num) + offline_scales[:, :head_num].view(-1, 1, head_num).expand(offline_scales.shape[0], att_batch_size, head_num) ) self.v_descale = ( - offline_scales[:, head_num:].view(-1, 1, head_num).expand(offline_scales.shape[0], batch_size, head_num) + offline_scales[:, head_num:].view(-1, 1, head_num).expand(offline_scales.shape[0], att_batch_size, head_num) ) return @@ -183,7 +185,7 @@ def _fp8_decode_att( cu_seqlens_q=self.cu_seqlens_q, cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, - causal=getattr(self.infer_state, "decode_causal", True), + causal=self.infer_state.decode_causal, window_size=(-1, -1), softcap=0.0, q_descale=q_scale.view(self.infer_state.batch_size, k_head_num), diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index 03c277c16c..c07cea8d4c 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -1,11 +1,10 @@ import dataclasses import torch from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl -from typing import Tuple +from typing import Optional, TYPE_CHECKING, Tuple from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.sgl_utils import flash_attn_with_kvcache -from lightllm.utils.envs_utils import enable_dynamic_mtp_verify, get_env_start_args -from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_mtp_fa3_decode_params, page_table_copy +from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_spec_fa3_decode_params, page_table_copy from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor from lightllm.utils.sgl_utils import flash_attn_varlen_func @@ -90,7 +89,7 @@ def _mla_prefill_att( max_seqlen_q=self.infer_state.max_q_seq_len, max_seqlen_k=self.infer_state.max_kv_seq_len, softmax_scale=softmax_scale, - causal=getattr(self.infer_state, "prefill_causal", True), + causal=self.infer_state.prefill_causal, return_softmax_lse=False, ) return o_tensor @@ -107,41 +106,40 @@ class MlaFa3DecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: MlaFa3AttBackend = self.backend - args_mtp_step = getattr(self.infer_state, "decode_mtp_step", None) - if args_mtp_step is None: - args_mtp_step = get_env_start_args().mtp_step - if self.infer_state.disable_mtp_decode_att: - args_mtp_step = 0 - is_block_draft_decode = getattr(self.infer_state, "is_draft_model", False) - is_dynamic_mtp = ( - args_mtp_step > 0 - and enable_dynamic_mtp_verify() - and not is_block_draft_decode - and not getattr(self.infer_state, "use_static_mtp_layout", False) - ) + draft_step = self.infer_state.draft_step + decode_rows_per_request = draft_step + 1 + uses_dynamic_spec_verify_layout = self.backend.uses_dynamic_spec_verify_layout(self.infer_state) + + if draft_step > 0 and not uses_dynamic_spec_verify_layout: + assert self.infer_state.batch_size % decode_rows_per_request == 0, ( + "FA3 fixed-layout decode requires batch_size to be divisible by draft_step + 1, " + f"got batch_size={self.infer_state.batch_size}, draft_step={draft_step}." + ) - if is_dynamic_mtp: - att_batch_size = self.infer_state.batch_size - (b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len,) = build_dynamic_mtp_fa3_decode_params( + # 修正 mtp 在 fa3 下的输入。 + if uses_dynamic_spec_verify_layout: + (b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len,) = build_dynamic_spec_fa3_decode_params( b_req_idx=self.infer_state.b_req_idx, b_seq_len=self.infer_state.b_seq_len, b_mark_shared_group=self.infer_state.b_mark_shared_group, - att_batch_size=att_batch_size, + att_batch_size=self.infer_state.batch_size, hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, ) - b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) - self.cu_seqlens_q = b1_cu_q_seq_len.int() - self.cu_seqlens_k = b1_cu_kv_seq_len.int() - elif args_mtp_step > 0: - # 修正 mtp 在 fa3 下的输入。 - mtp_size = args_mtp_step + 1 + elif draft_step > 0: b_q_seq_len = torch.full( - (self.infer_state.b_seq_len.shape[0] // mtp_size,), - fill_value=mtp_size, + (self.infer_state.b_seq_len.shape[0] // decode_rows_per_request,), + fill_value=decode_rows_per_request, dtype=torch.int32, device=self.infer_state.b_seq_len.device, ) - b_kv_seq_len = self.infer_state.b_seq_len[mtp_size - 1 :: mtp_size] + b_kv_seq_len = self.infer_state.b_seq_len[draft_step::decode_rows_per_request] + b_att_req_idx = self.infer_state.b_req_idx[draft_step::decode_rows_per_request] + self.b_att_seq_len = b_kv_seq_len.contiguous() + else: + b_att_req_idx = self.infer_state.b_req_idx + self.b_att_seq_len = self.infer_state.b_seq_len + + if draft_step > 0: b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) self.cu_seqlens_q = b1_cu_q_seq_len.int() self.cu_seqlens_k = b1_cu_kv_seq_len.int() @@ -149,10 +147,7 @@ def init_state(self): self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() - if not is_dynamic_mtp: - assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 - att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) - + att_batch_size = b_att_req_idx.shape[0] model = self.backend.model # 可以使用 cuda graph的时候从 buffer中申请 if ( @@ -170,29 +165,12 @@ def init_state(self): device=self.infer_state.input_ids.device, ) - if is_dynamic_mtp: - page_table_copy( - page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], - req_to_token_indexs=model.req_manager.req_to_token_indexs, - b_req_idx=b_att_req_idx, - ) - self.decode_max_q_seq_len = args_mtp_step + 1 - elif args_mtp_step > 0: - page_table_copy( - page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], - req_to_token_indexs=model.req_manager.req_to_token_indexs, - b_req_idx=self.infer_state.b_req_idx[args_mtp_step :: (args_mtp_step + 1)], - ) - self.b_att_seq_len = self.infer_state.b_seq_len[args_mtp_step :: (args_mtp_step + 1)].contiguous() - self.decode_max_q_seq_len = args_mtp_step + 1 - else: - page_table_copy( - page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], - req_to_token_indexs=model.req_manager.req_to_token_indexs, - b_req_idx=self.infer_state.b_req_idx, - ) - self.b_att_seq_len = self.infer_state.b_seq_len - self.decode_max_q_seq_len = 1 + page_table_copy( + page_table=self.page_table[:, : self.infer_state.max_kv_seq_len], + req_to_token_indexs=model.req_manager.req_to_token_indexs, + b_req_idx=b_att_req_idx, + ) + self.decode_max_q_seq_len = decode_rows_per_request return def copy_for_decode_cuda_graph(self, new_state: "MlaFa3DecodeAttState"): @@ -250,7 +228,7 @@ def _mla_decode_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, softmax_scale=softmax_scale, - causal=getattr(self.infer_state, "decode_causal", True), + causal=self.infer_state.decode_causal, window_size=(-1, -1), softcap=0.0, k_descale=k_descale, diff --git a/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index 482d240283..2441e1e52c 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -2,7 +2,7 @@ import torch from typing import TYPE_CHECKING from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl -from lightllm.utils.envs_utils import enable_dynamic_mtp_verify, get_env_start_args, get_llm_data_type +from lightllm.utils.envs_utils import get_env_start_args, get_llm_data_type from lightllm.common.basemodel.triton_kernel.linear_att.causal_conv1d import causal_conv1d_fn from lightllm.common.basemodel.triton_kernel.linear_att.fused_gdn_gating import fused_gdn_gating from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import chunk_gated_delta_rule @@ -10,6 +10,9 @@ from lightllm.common.basemodel.triton_kernel.linear_att.mtp_fused_recurrent import ( mtp_fused_recurrent_gated_delta_rule, ) +from lightllm.common.basemodel.triton_kernel.linear_att.spec_state_params import ( + build_dynamic_spec_linear_att_state_params, +) from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import fused_recurrent_gated_delta_rule if TYPE_CHECKING: @@ -26,7 +29,7 @@ def __init__(self, model: "TpPartBaseModel"): def _init_linear_layer_metadata(self, network_config, tp_world_size): - self.mtp_step = get_env_start_args().mtp_step + self.max_draft_step = get_env_start_args().mtp_step # Linear attention specific dimensions self.num_v_heads = network_config["linear_num_value_heads"] @@ -109,12 +112,12 @@ class LinearAttPrefillAttState(BasePrefillAttState): def init_state(self): backend: LinearAttBackend = self.backend - mtp_step = backend.mtp_step + max_draft_step = backend.max_draft_step # 每次 _prefill 都会在 runtime infer_state 上调用 init_state。 # prefill cuda graph 回调必须走 new_infer_state.prefill_att_state1, # 才能读到这里按当前 batch(含 token padding 后的 dummy request)更新的索引。 self.b_conv_buffer_idx = self.infer_state.b_req_idx - self.b_ssm_buffer_idx = self.infer_state.b_req_idx * (mtp_step + 1) + self.b_ssm_buffer_idx = self.infer_state.b_req_idx * (max_draft_step + 1) return def prefill_att( @@ -137,8 +140,8 @@ def prefill_att( conv_states, ssm_states = self.infer_state.req_manager.get_mamba_cache(layer_num) # 在开启了mtp的时候,conv 状态的最后一维可能存在冗余的部分,需要进行切片对齐。 # prefill 模式下,使用不到这几个维度,所以需要扣除掉, - if backend.mtp_step > 0: - conv_states = conv_states[:, :, : -backend.mtp_step] + if backend.max_draft_step > 0: + conv_states = conv_states[:, :, : -backend.max_draft_step] mixed_qkv, z, b, a = backend._split_qkvzba(mixed_qkvzba) core_attn_out = self._gdn_prefill_kernel( mixed_qkv, conv_states, ssm_states, a, b, self.infer_state, layer_weight @@ -199,85 +202,46 @@ class LinearAttDecodeAttState(BaseDecodeAttState): b_conv_buffer_idx: torch.Tensor = None b_ssm_buffer_idx: torch.Tensor = None - b1_mtp_cu_q_seq_len: torch.Tensor = None + b1_spec_cu_q_seq_len: torch.Tensor = None b_num_accepted_tokens: torch.Tensor = None - def _uses_dynamic_mtp_layout(self) -> bool: - return ( - self.backend.mtp_step > 0 - and enable_dynamic_mtp_verify() - and not getattr(self.infer_state, "use_static_mtp_layout", False) - ) - - def prepare_for_forward(self): - """Build compact GDN row metadata as part of the captured forward. - - CUDA graph replay copies the primary ModelInput tensors into the graph - state. Deriving these tensors inside the graph makes every replay use - the current compact request layout instead of capture-time metadata. - """ - - if not self._uses_dynamic_mtp_layout(): - return - - from lightllm.common.basemodel.triton_kernel.linear_att.mtp_state_params import ( - build_dynamic_mtp_linear_att_state_params, - ) - - backend: LinearAttBackend = self.backend - batch_size = self.infer_state.batch_size - ( - self.b1_mtp_cu_q_seq_len, - self.b_conv_buffer_idx, - self.b_num_accepted_tokens, - ) = build_dynamic_mtp_linear_att_state_params( - b_req_idx=self.infer_state.b_req_idx, - b_mtp_index=self.infer_state.b_mtp_index, - req_to_mtp_state_index=self.infer_state.req_manager.req_to_mtp_state_index, - hold_req_id=self.infer_state.req_manager.HOLD_REQUEST_ID, - ) - self.b_ssm_buffer_idx = self.b_conv_buffer_idx.view(batch_size, 1) * (backend.mtp_step + 1) + torch.arange( - backend.mtp_step + 1, - device=self.infer_state.b_req_idx.device, - dtype=self.infer_state.b_req_idx.dtype, - ).view(1, backend.mtp_step + 1) - return - def init_state(self): - backend: LinearAttBackend = self.backend - mtp_step = backend.mtp_step + draft_step = self.infer_state.draft_step - # decode 模式下 - if mtp_step == 0: - # 非mtp模式下,不需要额外状态 + if draft_step == 0: self.b_conv_buffer_idx = self.infer_state.b_req_idx self.b_ssm_buffer_idx = self.infer_state.b_req_idx return - if mtp_step > 0: - # mtp 模式下 - batch_size = self.infer_state.batch_size - if self._uses_dynamic_mtp_layout(): - # This must run inside _token_forward so CUDA graph replay - # recomputes it from the current compact row layout. - return - - att_batch_size = batch_size // (mtp_step + 1) - assert batch_size % (mtp_step + 1) == 0 - - device = self.infer_state.b_req_idx.device - - # shape 为 [att_batch_size + 1] - self.b1_mtp_cu_q_seq_len = torch.arange(0, batch_size + 1, mtp_step + 1, dtype=torch.int32, device=device) - # shape 为 [att_batch_size] - self.b_conv_buffer_idx = self.infer_state.b_req_idx.view(att_batch_size, mtp_step + 1)[:, 0].contiguous() - self.b_ssm_buffer_idx = (self.b_conv_buffer_idx * (mtp_step + 1)).view(att_batch_size, 1) + torch.arange( - mtp_step + 1, device=device, dtype=self.infer_state.b_req_idx.dtype - ).view(1, mtp_step + 1) - # shape 为 [att_batch_size] - # 上一步接受的数量,用于linear att 的decode mtp 算子定位正确的conv 和 ssm信息的起点。 + batch_size = self.infer_state.batch_size + device = self.infer_state.b_req_idx.device + if self.backend.uses_dynamic_spec_verify_layout(self.infer_state): + ( + self.b1_spec_cu_q_seq_len, + self.b_conv_buffer_idx, + self.b_num_accepted_tokens, + ) = build_dynamic_spec_linear_att_state_params( + b_req_idx=self.infer_state.b_req_idx, + b_mtp_index=self.infer_state.b_mtp_index, + req_to_mtp_state_index=self.infer_state.req_manager.req_to_mtp_state_index, + hold_req_id=self.infer_state.req_manager.HOLD_REQUEST_ID, + ) + else: + assert batch_size % (draft_step + 1) == 0, ( + "GDN fixed-layout decode requires batch_size to be divisible by draft_step + 1, " + f"got batch_size={batch_size}, draft_step={draft_step}." + ) + att_batch_size = batch_size // (draft_step + 1) + self.b1_spec_cu_q_seq_len = torch.arange( + 0, batch_size + 1, draft_step + 1, dtype=torch.int32, device=device + ) + self.b_conv_buffer_idx = self.infer_state.b_req_idx.view(att_batch_size, draft_step + 1)[:, 0].contiguous() self.b_num_accepted_tokens = self.infer_state.req_manager.req_to_mtp_state_index[self.b_conv_buffer_idx] + 1 - return + + # Each request owns one recurrent-state slot per verify row. + state_offsets = torch.arange(draft_step + 1, device=device, dtype=self.infer_state.b_req_idx.dtype) + self.b_ssm_buffer_idx = self.b_conv_buffer_idx[:, None] * (draft_step + 1) + state_offsets[None, :] + return def decode_att( self, @@ -299,9 +263,8 @@ def decode_att( mixed_qkv, z, b, a = backend._split_qkvzba(mixed_qkvzba) conv_states, ssm_states = self.infer_state.req_manager.get_mamba_cache(layer_num) - if backend.mtp_step > 0: - # MTP 模式下,使用线性层 MTP 状态。 - core_attn_out = self._gdn_mtp_kernel( + if self.infer_state.draft_step > 0: + core_attn_out = self._gdn_spec_kernel( mixed_qkv, conv_states, ssm_states, @@ -371,7 +334,7 @@ def _gdn_decode_kernel( ) return core_attn_out, z - def _gdn_mtp_kernel( + def _gdn_spec_kernel( self, mixed_qkv: torch.Tensor, conv_states: torch.Tensor, @@ -387,12 +350,12 @@ def _gdn_mtp_kernel( backend: LinearAttBackend = self.backend - cu_seqlens_q = self.b1_mtp_cu_q_seq_len + cu_seqlens_q = self.b1_spec_cu_q_seq_len mixed_qkv = causal_conv1d_update_spec( mixed_qkv, conv_states, layer_weight.linear_conv1d.mm_param.weight, - mtp_step=backend.mtp_step, + mtp_step=infer_state.draft_step, bias=layer_weight.linear_conv1d.bias, activation=backend.activation, conv_state_indices=self.b_conv_buffer_idx, diff --git a/lightllm/common/basemodel/attention/triton/fp.py b/lightllm/common/basemodel/attention/triton/fp.py index 9e1407b483..f2c1127fbb 100644 --- a/lightllm/common/basemodel/attention/triton/fp.py +++ b/lightllm/common/basemodel/attention/triton/fp.py @@ -1,8 +1,7 @@ import dataclasses import torch - -from lightllm.utils.envs_utils import enable_dynamic_mtp_verify, get_env_start_args, enable_triton_mtp_kernel from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl +from typing import Optional class TritonAttBackend(BaseAttBackend): @@ -94,21 +93,8 @@ def _nomarl_prefill_att( @dataclasses.dataclass class TritonDecodeAttState(BaseDecodeAttState): - # MTP related state variables - b_mark_shared_group: torch.Tensor = None - def init_state(self): - args_mtp_step = getattr(self.infer_state, "decode_mtp_step", None) - if args_mtp_step is None: - args_mtp_step = get_env_start_args().mtp_step - if self.infer_state.disable_mtp_decode_att: - args_mtp_step = 0 - - if args_mtp_step > 0: - # MTP mode initialization - self.b_mark_shared_group = self.infer_state.b_mark_shared_group - else: - self.b_mark_shared_group = None + pass def copy_for_decode_cuda_graph(self, new_state: "TritonDecodeAttState"): super().copy_for_decode_cuda_graph(new_state) @@ -126,20 +112,16 @@ def decode_att( assert att_control.tp_alibi is not None return self._alibi_decode_att(q=q, k=k, v=v, att_control=att_control, alloc_func=alloc_func) else: - args_mtp_step = getattr(self.infer_state, "decode_mtp_step", None) - if args_mtp_step is None: - args_mtp_step = get_env_start_args().mtp_step - if self.infer_state.disable_mtp_decode_att: - args_mtp_step = 0 + draft_step = self.infer_state.draft_step q_head_num = q.shape[1] k_head_num = k.shape[1] - if args_mtp_step > 0 and (enable_dynamic_mtp_verify() or enable_triton_mtp_kernel()): - # MTP mode: use mtp diverse attention - assert q_head_num >= k_head_num, "MTP diverse attention requires q_head_num >= k_head_num" - return self._dynamic_mtp_decode_gqa_att(q=q, k=k, v=v, alloc_func=alloc_func) + if draft_step > 0: + assert q_head_num >= k_head_num, "speculative decode requires q_head_num >= k_head_num" + return self._spec_decode_gqa_att(q=q, k=k, v=v, alloc_func=alloc_func) elif q_head_num == k_head_num: + assert att_control.use_sliding_window is False, "sliding_window not supported in non-gqa attention yet" return self._normal_decode_flash_decoding_att(q=q, k=k, v=v, alloc_func=alloc_func) elif q_head_num > k_head_num: return self._normal_decode_gqa_flash_decoding_att( @@ -229,7 +211,7 @@ def _normal_decode_gqa_flash_decoding_att( return out - def _dynamic_mtp_decode_gqa_att( + def _spec_decode_gqa_att( self, q: torch.Tensor, k: torch.Tensor, @@ -240,18 +222,14 @@ def _dynamic_mtp_decode_gqa_att( token_decode_attention_mtp_diverse_single_token, ) - b_seq_len = self.infer_state.b_seq_len - # 在动态 MTP 验证模式下,使用 infer_state.b_mark_shared_group(从 model_input 传递) - # 在静态 MTP 模式下,使用 self.b_mark_shared_group(在 init_state 中初始化) - b_mark_shared_group = self.infer_state.b_mark_shared_group out = token_decode_attention_mtp_diverse_single_token( q=q, k=k, v=v, Req_to_tokens=self.infer_state.req_manager.req_to_token_indexs, B_req_idx=self.infer_state.b_req_idx, - b_seq_len=b_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_seq_len=self.infer_state.b_seq_len, + b_mark_shared_group=self.infer_state.b_mark_shared_group, alloc_tensor_func=alloc_func, ) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 7dcda28e41..85ca25cfd6 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -14,6 +14,7 @@ from lightllm.common.basemodel.infer_struct import InferStateInfo from lightllm.common.kv_cache_mem_manager import MemoryManager from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class +from lightllm.common.req_manager import ReqManager from lightllm.common.infer_utils import init_req_to_token_indexes from lightllm.common.build_utils import repair_config from lightllm.common.basemodel.triton_kernel.copy_kv_index_to_req import copy_kv_index_to_req @@ -24,20 +25,20 @@ from lightllm.common.basemodel.triton_kernel.gather_token_id import gather_token, gather_token_prefill_decode_mixed from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_dp_world_size -from lightllm.utils.envs_utils import ( - enable_triton_mtp_kernel, - get_env_start_args, - get_llm_data_type, - get_added_mtp_kv_layer_num, -) +from lightllm.utils.envs_utils import get_env_start_args, get_llm_data_type, get_added_mtp_kv_layer_num from lightllm.distributed.communication_op import dist_group_manager from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.hidden_collector import ( + HiddenCollector, + unpad_collected_hidden, +) from lightllm.utils.custom_kernel_utis import pad2dim_tensor_to_new_batch from lightllm.utils.envs_utils import ( set_model_init_status, enable_diverse_mode_gqa_decode_fast_kernel, enable_full_att_decode_tune, - enable_dynamic_mtp_verify, + enable_dynamic_spec, + enable_triton_mtp_kernel, ) from lightllm.common.triton_utils.autotuner import Autotuner from lightllm.utils.infer_utils import post_empty_cache @@ -54,6 +55,8 @@ class TpPartBaseModel: + is_mtp_draft_model = False + # weight class pre_and_post_weight_class = None transformer_weight_class = None @@ -83,27 +86,14 @@ def __init__(self, kvargs): self.return_all_prompt_logics = kvargs.get("return_all_prompt_logics", False) assert not (self.is_token_healing and self.return_all_prompt_logics), "can not be true in same time" self.data_type = get_llm_data_type() - mtp_step = get_env_start_args().mtp_step - # A graph for more logical requests than the scheduler can ever run is - # unreachable. In MTP modes the value is expanded by ``mtp_step + 1`` - # below, so leaving the CLI default (often 256) uncapped can otherwise - # create multi-GiB full-vocabulary graph outputs for a server limited - # to only a few dozen concurrent requests. - self.graph_max_batch_size = min( - kvargs.get("graph_max_batch_size", 16), - self.max_req_num, - ) + self.decode_batch_multiplier = kvargs.get("decode_batch_multiplier", 1) + self.graph_max_batch_size = kvargs.get("graph_max_batch_size", 16) self.graph_max_batch_size = ( self.graph_max_batch_size // 2 if get_env_start_args().enable_decode_microbatch_overlap else self.graph_max_batch_size ) - # mtp 模式下需要修缮对应的最大batch size,为 (mtp_step + 1) 的倍数 - self.graph_max_batch_size = self.graph_max_batch_size * (mtp_step + 1) - self.graph_split_batch_size = int(kvargs.get("graph_split_batch_size", self.args.graph_split_batch_size)) - self.graph_grow_step_size = int(kvargs.get("graph_grow_step_size", self.args.graph_grow_step_size)) - assert self.graph_split_batch_size > 0 - assert self.graph_grow_step_size > 0 + self.graph_max_batch_size = self.graph_max_batch_size * self.decode_batch_multiplier self.graph_max_len_in_batch = kvargs.get("graph_max_len_in_batch", 8192) self.disable_cudagraph = kvargs.get("disable_cudagraph", False) @@ -116,10 +106,8 @@ def __init__(self, kvargs): self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) self.prefill_graph: PrefillCudaGraph = None - self.spec_adapter = None self._init_config() - self._verify_must() self._verify_params() self._init_quant() @@ -149,6 +137,7 @@ def __init__(self, kvargs): logger.info(f"use prefill att backend1: {self.prefill_att_backend1.__class__.__name__}") logger.info(f"use decode att backend1: {self.decode_att_backend1.__class__.__name__}") + self._init_hidden_collector(hidden_layer_ids=kvargs.get("hidden_layer_ids")) self._autotune_warmup() self._full_att_decode_autotune() self._init_padded_req() @@ -159,26 +148,6 @@ def __init__(self, kvargs): set_model_init_status(True) return - def _wait_other_modules_ready(self): - for event in self.wait_events: - event.wait() - return - - def set_spec_adapter(self, spec_adapter): - self.spec_adapter = spec_adapter - if self.graph is not None: - self.graph.set_spec_adapter(spec_adapter, model=self) - # Models build their initial decode graphs before the speculative - # runtime exists. Attaching the adapter changes both graph batch - # sizes and cache keys, so eagerly capture the speculative graph - # variants on dummy state now. Lazy capture during the first live - # request can execute in-place KV/linear-state updates repeatedly. - if get_env_start_args().enable_decode_microbatch_overlap: - self.graph.warmup_overlap(self) - else: - self.graph.warmup(self) - return - def _init_config(self): with open(os.path.join(self.weight_dir_, "config.json"), "r") as json_file: self.config = json.load(json_file) @@ -188,12 +157,6 @@ def _init_config(self): repair_config(self.config, same_names=["num_hidden_layers", "n_layer"]) if self.finetune_config: self.config["vocab_size"] = self.finetune_config.vocab_size - - # eagle3 mode 下,需要修改 vocab_size 为 draft_vocab_size, 其他场景 - # 这个代码并不会生效。 - if "draft_vocab_size" in self.config.keys(): - self.config["target_vocab_size"] = self.config["vocab_size"] - self.config["vocab_size"] = self.config["draft_vocab_size"] return @final @@ -270,8 +233,6 @@ def _check_mem_size(self): return def _init_req_manager(self): - from lightllm.common.req_manager import ReqManager - create_max_seq_len = 0 if self.batch_max_tokens is not None: @@ -313,6 +274,7 @@ def _init_att_backend1(self): return def _init_cudagraph(self): + batch_multiplier = 1 if enable_dynamic_spec() and not self.is_mtp_draft_model else self.decode_batch_multiplier self.graph = ( None if self.disable_cudagraph @@ -320,8 +282,8 @@ def _init_cudagraph(self): max_batch_size=self.graph_max_batch_size, max_len_in_batch=self.graph_max_len_in_batch, tp_world_size=self.tp_world_size_, - graph_split_batch_size=self.graph_split_batch_size, - graph_grow_step_size=self.graph_grow_step_size, + batch_multiplier=batch_multiplier, + capture_infer_cost=enable_dynamic_spec(), ) ) if self.graph is not None: @@ -361,7 +323,7 @@ def _full_att_decode_autotune(self): if self.disable_cudagraph: return # Only tune on the main model; MTP draft models skip this path. - if getattr(self, "is_mtp_draft_model", False): + if self.is_mtp_draft_model: return # Opt-in switch for FA3 full-attention decode num_splits tuning. @@ -372,7 +334,7 @@ def _full_att_decode_autotune(self): # Only Fa3AttBackend decode path needs this num_splits warmup. decode_backends = [ self.decode_att_backend, - getattr(self, "decode_att_backend1", None), + self.decode_att_backend1, ] if not any( backend is not None and backend.__class__.__name__ == "Fa3AttBackend" for backend in decode_backends @@ -381,23 +343,36 @@ def _full_att_decode_autotune(self): from lightllm.utils.sgl_utils import fa3_decode_autotune + batch_multiplier = 1 if enable_dynamic_spec() and not self.is_mtp_draft_model else self.decode_batch_multiplier cuda_graph_batch_sizes = CudaGraph.gen_cuda_graph_batch_sizes( max_batch_size=self.graph_max_batch_size, tp_world_size=self.tp_world_size_, + batch_multiplier=batch_multiplier, ) - fa3_decode_autotune(self, cuda_graph_batch_sizes) + fa3_decode_autotune(self, cuda_graph_batch_sizes, batch_multiplier=batch_multiplier) return def _init_custom(self): pass + def _init_hidden_collector(self, hidden_layer_ids): + microbatch_count = ( + 2 if self.args.enable_prefill_microbatch_overlap or self.args.enable_decode_microbatch_overlap else 1 + ) + self.hidden_collector = HiddenCollector( + model=self, + spec_mode=self.args.mtp_mode, + layer_ids=hidden_layer_ids, + microbatch_count=microbatch_count, + ) + @torch.no_grad() def forward(self, model_input: ModelInput): model_input.to_cuda() assert model_input.mem_indexes.is_cuda if model_input.is_prefill: - return self._prefill(model_input) + return self._prefill(model_input=model_input) else: return self._decode(model_input) @@ -427,7 +402,7 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) if enable_diverse_mode_gqa_decode_fast_kernel(): infer_state.b_shared_seq_len = model_input.b_shared_seq_len infer_state.b_mark_shared_group = model_input.b_mark_shared_group - elif enable_dynamic_mtp_verify() or enable_triton_mtp_kernel(): + elif enable_dynamic_spec() or enable_triton_mtp_kernel(): infer_state.b_mark_shared_group = model_input.b_mark_shared_group infer_state.multimodal_params = model_input.multimodal_params @@ -441,8 +416,7 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) # 特殊模型,特殊模式的特定变量初始化操作。 infer_state.mtp_draft_input_hiddens = model_input.mtp_draft_input_hiddens - infer_state.disable_mtp_decode_att = model_input.disable_mtp_decode_att - infer_state.use_static_mtp_layout = model_input.use_static_mtp_layout + infer_state.draft_step = model_input.draft_step if infer_state.is_prefill: infer_state.prefill_att_state = self.prefill_att_backend.create_att_prefill_state(infer_state=infer_state) @@ -501,7 +475,7 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.b_mark_shared_group = F.pad( new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=1 ) - elif enable_dynamic_mtp_verify() or enable_triton_mtp_kernel(): + elif enable_dynamic_spec() or enable_triton_mtp_kernel(): assert new_model_input.b_mark_shared_group is not None new_model_input.b_mark_shared_group = F.pad( new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=0 @@ -564,60 +538,27 @@ def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle new_model_input.check_input() return new_model_input - def _create_unpad_decode_model_output( - self, - model_output: ModelOutput, - origin_batch_size: int, - microbatch_index: int = 0, - ): + def _create_unpad_decode_model_output(self, model_output: ModelOutput, origin_batch_size: int): padded_batch_size = model_output.logits.shape[0] if padded_batch_size == origin_batch_size: return model_output new_model_output = copy.copy(model_output) new_model_output.logits = new_model_output.logits[0:origin_batch_size] - if new_model_output.mtp_draft_token_ids is not None: - new_model_output.mtp_draft_token_ids = new_model_output.mtp_draft_token_ids[0:origin_batch_size] - if new_model_output.mtp_draft_confidence_logits is not None: - confidence_rows = new_model_output.mtp_draft_confidence_logits.shape[0] - if confidence_rows == padded_batch_size: - confidence_origin_rows = origin_batch_size - else: - assert padded_batch_size % confidence_rows == 0, ( - "padded decode logits rows must be divisible by confidence rows: " - f"{padded_batch_size}, got {confidence_rows}" - ) - rows_per_confidence = padded_batch_size // confidence_rows - assert origin_batch_size % rows_per_confidence == 0, ( - "origin decode rows must align with confidence row grouping: " - f"{origin_batch_size}, rows_per_confidence={rows_per_confidence}" - ) - confidence_origin_rows = origin_batch_size // rows_per_confidence - new_model_output.mtp_draft_confidence_logits = new_model_output.mtp_draft_confidence_logits[ - 0:confidence_origin_rows - ] - if self.spec_adapter is not None: - self.spec_adapter.unpad_hidden(token_num=origin_batch_size, microbatch_index=microbatch_index) - + new_model_output.spec_hidden = unpad_collected_hidden(new_model_output.spec_hidden, origin_batch_size) return new_model_output def _create_unpad_prefill_model_output( - self, - padded_model_output: ModelOutput, - origin_handle_token_num: int, - origin_batch_size: int, - microbatch_index: int = 0, + self, padded_model_output: ModelOutput, origin_handle_token_num: int, origin_batch_size: int ): new_model_output = copy.copy(padded_model_output) # logits 始终只对应每个请求最后一个位置,移除 padding 的 req 对应的行。 new_model_output.logits = new_model_output.logits[0:origin_batch_size] + new_model_output.spec_hidden = unpad_collected_hidden(new_model_output.spec_hidden, origin_handle_token_num) # prompt_logics 保存整个 prefill 阶段所有 token 位置的 logits, # 按实际处理的 token 数量裁剪掉 padding 部分(仅 return_all_prompt_logics 模式下非空)。 if new_model_output.prompt_logics is not None: new_model_output.prompt_logics = new_model_output.prompt_logics[0:origin_handle_token_num] - if self.spec_adapter is not None: - self.spec_adapter.unpad_hidden(token_num=origin_handle_token_num, microbatch_index=microbatch_index) - return new_model_output def _prefill( @@ -666,7 +607,7 @@ def _prefill( infer_state.init_some_extra_state(self) infer_state.init_att_state() - model_output = self._context_forward(infer_state) + model_output = self._context_forward(infer_state=infer_state) model_output = self._create_unpad_prefill_model_output( padded_model_output=model_output, @@ -695,15 +636,14 @@ def _decode( infer_batch_size = model_input.batch_size if self.graph is not None and self.graph.can_run( - batch_size=infer_batch_size, - max_len_in_batch=model_input.max_kv_seq_len, + batch_size=infer_batch_size, max_len_in_batch=model_input.max_kv_seq_len ): infer_batch_size = self.graph.find_closest_graph_batch_size(batch_size=infer_batch_size) model_input = self._create_padded_decode_model_input( model_input=model_input, new_batch_size=infer_batch_size ) - need_capture = self.graph.need_capture(infer_batch_size, model_context=model_input) infer_state = self._create_inferstate(model_input) + need_capture = self.graph.need_capture(infer_batch_size) infer_state.is_cuda_graph = need_capture copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -749,26 +689,19 @@ def _context_forward(self, infer_state: InferStateInfo): input_embs = self.pre_infer._tpsp_sp_split(input=input_embs, infer_state=infer_state) input_tensors = [input_embs] + hidden_collector = HiddenCollector() if Autotuner.is_autotune_warmup() else self.hidden_collector def prefill_func(input_tensors, infer_state): - spec_context = ( - self.spec_adapter.create_forward_context(self, infer_state) if self.spec_adapter is not None else None - ) _input_embs = input_tensors[0] for i in range(self.layers_num): layer = self.layers_infer[i] _input_embs = layer.context_forward(_input_embs, infer_state, self.trans_layers_weight[i]) - if spec_context is not None: - spec_context.add_hidden( - layer_index=i, - layer_num=self.layers_num, - hidden=_input_embs, - ) - - layer_hidden = spec_context.build_layer_hidden() if spec_context is not None else None - if layer_hidden is not None: - return [_input_embs, layer_hidden] - return [_input_embs] + hidden_collector.add( + layer_index=i, + hidden=_input_embs, + ) + + return hidden_collector.prefill_outputs(_input_embs) handle_token_num = infer_state.input_ids.shape[0] @@ -800,80 +733,43 @@ def prefill_func(input_tensors, infer_state): last_input_embs = infer_state._all_to_all_unbalance_get(data=last_input_embs) predict_logits = self.post_infer.token_forward(last_input_embs, infer_state, self.pre_post_weight) - model_output = ModelOutput(logits=predict_logits, prompt_logics=infer_state.prompt_logics) - - if self.spec_adapter is not None: - if len(output_tensors) > 1: - spec_hidden = self.pre_infer._tpsp_allgather(input=output_tensors[1], infer_state=infer_state) - if infer_state.need_dp_prefill_balance: - spec_hidden = infer_state._all_to_all_unbalance_get(data=spec_hidden) - else: - spec_hidden = last_input_embs - self.spec_adapter.capture_hidden( - infer_state=infer_state, - hidden=spec_hidden.contiguous(), - final_hidden=last_input_embs.contiguous(), - ) + spec_hidden = hidden_collector.finish( + infer_state=infer_state, + final_hidden=last_input_embs, + forward_outputs=output_tensors, + ) + model_output = ModelOutput( + logits=predict_logits, + spec_hidden=spec_hidden, + prompt_logics=infer_state.prompt_logics, + ) # 在开启使用deepep的时候,需要调用clear_deepep_buffer做资源清理,没有启用的时候 # 该调用没有实际意义 dist_group_manager.clear_deepep_buffer() return model_output - @final def _token_forward(self, infer_state: InferStateInfo): - # Some derived decode metadata depends on runtime tensor values and - # must therefore be computed inside CUDA graph capture/replay. The - # attention states cache it for reuse by every transformer layer. - infer_state.decode_att_state.prepare_for_forward() - if infer_state.decode_att_state1 is not None: - infer_state.decode_att_state1.prepare_for_forward() - input_ids = infer_state.input_ids cuda_input_ids = input_ids input_embs = self.pre_infer.token_forward(cuda_input_ids, infer_state, self.pre_post_weight) input_embs = self.pre_infer._tpsp_sp_split(input=input_embs, infer_state=infer_state) - spec_context = ( - self.spec_adapter.create_forward_context(self, infer_state) if self.spec_adapter is not None else None - ) for i in range(self.layers_num): layer = self.layers_infer[i] input_embs: torch.Tensor = layer.token_forward(input_embs, infer_state, self.trans_layers_weight[i]) - if spec_context is not None: - spec_context.add_hidden(layer_index=i, layer_num=self.layers_num, hidden=input_embs) + self.hidden_collector.add(layer_index=i, hidden=input_embs) last_input_embs = self.post_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) predict_logits: torch.Tensor = self.post_infer.token_forward( last_input_embs, infer_state=infer_state, layer_weight=self.pre_post_weight ) - mtp_draft_confidence_logits = None - pop_confidence_logits = getattr(self.post_infer, "pop_mtp_draft_confidence_logits", None) - if pop_confidence_logits is not None: - mtp_draft_confidence_logits = pop_confidence_logits() - - mtp_draft_token_ids = None - pop_draft_token_ids = getattr(self.post_infer, "pop_mtp_draft_token_ids", None) - if pop_draft_token_ids is not None: - mtp_draft_token_ids = pop_draft_token_ids() - - model_output = ModelOutput( - logits=predict_logits.contiguous(), - mtp_draft_confidence_logits=mtp_draft_confidence_logits, - mtp_draft_token_ids=mtp_draft_token_ids, + spec_hidden = self.hidden_collector.finish( + infer_state=infer_state, + final_hidden=last_input_embs, ) - - if spec_context is not None: - spec_hidden = spec_context.build_layer_hidden() - if spec_hidden is not None: - spec_hidden = self.pre_infer._tpsp_allgather(input=spec_hidden, infer_state=infer_state) - else: - spec_hidden = last_input_embs - spec_context.capture( - hidden=spec_hidden.contiguous(), - final_hidden=last_input_embs.contiguous(), - ) + model_output = ModelOutput(logits=predict_logits.contiguous(), spec_hidden=spec_hidden) # 在 cuda graph 模式下,输出需要转为 no ref tensor, 加强mem pool 的复用,降低显存的使用。 if infer_state.is_cuda_graph: @@ -959,13 +855,11 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod padded_model_output=model_output0, origin_handle_token_num=origin_handle_token_num0, origin_batch_size=origin_batch_size0, - microbatch_index=0, ) model_output1 = self._create_unpad_prefill_model_output( padded_model_output=model_output1, origin_handle_token_num=origin_handle_token_num1, origin_batch_size=origin_batch_size1, - microbatch_index=1, ) # 在开启使用deepep的时候,需要调用clear_deepep_buffer做资源清理,没有启用的时候 # 该调用没有实际意义 @@ -1003,17 +897,12 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode if self.graph is not None and self.graph.can_run(infer_batch_size, max_len_in_batch): infer_batch_size = self.graph.find_closest_graph_batch_size(infer_batch_size) + need_capture = self.graph.need_capture(infer_batch_size) # TODO 如果支持动态步数的 mtp,在不同的mtp步上,model_input0 和 model_input1 的内部batch size可能不 # 一致,需要按照较高 batch size 进行graph的寻找,同时,进行有效的恢复。 padded_model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) padded_model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) - need_capture = self.graph.need_capture( - infer_batch_size, - model_context=padded_model_input0, - model_context1=padded_model_input1, - ) infer_state0 = self._create_inferstate(padded_model_input0, 0) - infer_state1 = self._create_inferstate(padded_model_input1, 1) infer_state0.is_cuda_graph = need_capture copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -1024,6 +913,7 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state0.init_some_extra_state(self) infer_state0.init_att_state() + infer_state1 = self._create_inferstate(padded_model_input1, 1) infer_state1.is_cuda_graph = need_capture copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -1047,12 +937,8 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode ) # TODO 动态 mtp fix - model_output0 = self._create_unpad_decode_model_output( - model_output0, origin_batch_size=origin_batch_size, microbatch_index=0 - ) - model_output1 = self._create_unpad_decode_model_output( - model_output1, origin_batch_size=origin_batch_size, microbatch_index=1 - ) + model_output0 = self._create_unpad_decode_model_output(model_output0, origin_batch_size=origin_batch_size) + model_output1 = self._create_unpad_decode_model_output(model_output1, origin_batch_size=origin_batch_size) else: model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) @@ -1077,12 +963,8 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state1.init_att_state() model_output0, model_output1 = self._overlap_tpsp_token_forward(infer_state0, infer_state1=infer_state1) - model_output0 = self._create_unpad_decode_model_output( - model_output0, origin_batch_size=origin_batch_size, microbatch_index=0 - ) - model_output1 = self._create_unpad_decode_model_output( - model_output1, origin_batch_size=origin_batch_size, microbatch_index=1 - ) + model_output0 = self._create_unpad_decode_model_output(model_output0, origin_batch_size=origin_batch_size) + model_output1 = self._create_unpad_decode_model_output(model_output1, origin_batch_size=origin_batch_size) return model_output0, model_output1 @@ -1109,6 +991,15 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state input_embs, input_embs1 = self.layers_infer[i].overlap_tpsp_context_forward( input_embs, input_embs1, infer_state, infer_state1, self.trans_layers_weight[i] ) + self.hidden_collector.add( + layer_index=i, + hidden=input_embs, + ) + self.hidden_collector.add( + layer_index=i, + hidden=input_embs1, + microbatch_index=1, + ) # 折叠模式调用完infer_state 和 infer_state1 上的hook函数后,input_embs 和 input_embs1 才具备正确的运算数据。 infer_state.call_overlap_hook() @@ -1125,22 +1016,25 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state ) g_cache_manager.cache_env_out() - model_output = ModelOutput(logits=predict_logits.contiguous(), prompt_logics=infer_state.prompt_logics) - model_output1 = ModelOutput(logits=predict_logits1.contiguous(), prompt_logics=infer_state1.prompt_logics) - - if self.spec_adapter is not None: - spec_context = self.spec_adapter.create_forward_context(self, infer_state) - input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) - if infer_state.need_dp_prefill_balance: - input_embs = infer_state._all_to_all_unbalance_get(data=input_embs) - spec_context.capture_final_hidden(input_embs.contiguous()) - - if self.spec_adapter is not None: - spec_context1 = self.spec_adapter.create_forward_context(self, infer_state1) - input_embs1 = self.pre_infer._tpsp_allgather(input=input_embs1, infer_state=infer_state1) - if infer_state1.need_dp_prefill_balance: - input_embs1 = infer_state1._all_to_all_unbalance_get(data=input_embs1) - spec_context1.capture_final_hidden(input_embs1.contiguous()) + spec_hidden = self.hidden_collector.finish( + infer_state=infer_state, + final_hidden=last_input_embs, + ) + spec_hidden1 = self.hidden_collector.finish( + infer_state=infer_state1, + final_hidden=last_input_embs1, + microbatch_index=1, + ) + model_output = ModelOutput( + logits=predict_logits.contiguous(), + spec_hidden=spec_hidden, + prompt_logics=infer_state.prompt_logics, + ) + model_output1 = ModelOutput( + logits=predict_logits1.contiguous(), + spec_hidden=spec_hidden1, + prompt_logics=infer_state1.prompt_logics, + ) return model_output, model_output1 @@ -1156,6 +1050,15 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: input_embs, input_embs1 = self.layers_infer[i].overlap_tpsp_token_forward( input_embs, input_embs1, infer_state, infer_state1, self.trans_layers_weight[i] ) + self.hidden_collector.add( + layer_index=i, + hidden=input_embs, + ) + self.hidden_collector.add( + layer_index=i, + hidden=input_embs1, + microbatch_index=1, + ) # 折叠模式调用完infer_state 上的hook函数后,input_embs 和 input_embs 才具备正确的运算数据。 infer_state.call_overlap_hook() @@ -1168,18 +1071,17 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: last_input_embs, last_input_embs1, infer_state, infer_state1, self.pre_post_weight ) - model_output = ModelOutput(logits=predict_logits.contiguous()) - model_output1 = ModelOutput(logits=predict_logits1.contiguous()) - - if self.spec_adapter is not None: - spec_context = self.spec_adapter.create_forward_context(self, infer_state) - input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) - spec_context.capture_final_hidden(input_embs.contiguous()) - - if self.spec_adapter is not None: - spec_context1 = self.spec_adapter.create_forward_context(self, infer_state1) - input_embs1 = self.pre_infer._tpsp_allgather(input=input_embs1, infer_state=infer_state1) - spec_context1.capture_final_hidden(input_embs1.contiguous()) + spec_hidden = self.hidden_collector.finish( + infer_state=infer_state, + final_hidden=last_input_embs, + ) + spec_hidden1 = self.hidden_collector.finish( + infer_state=infer_state1, + final_hidden=last_input_embs1, + microbatch_index=1, + ) + model_output = ModelOutput(logits=predict_logits.contiguous(), spec_hidden=spec_hidden) + model_output1 = ModelOutput(logits=predict_logits1.contiguous(), spec_hidden=spec_hidden1) if infer_state.is_cuda_graph: model_output.to_no_ref_tensor() @@ -1308,9 +1210,9 @@ def _autotune_warmup(self): multimodal_params=[{"images": [], "audios": []}], **self._gen_special_model_input(total_token_num), ) - model_output = self.forward( - model_input, - ) + model_input.to_cuda() + assert model_input.mem_indexes.is_cuda + model_output = self._prefill(model_input=model_input) del model_output self.req_manager.free_all() self.mem_manager.free_all() @@ -1391,8 +1293,7 @@ def _init_padded_req(self): def _gen_special_model_input(self, token_num: int): special_model_input = {} - is_mtp_draft_model = getattr(self, "is_mtp_draft_model", False) - if is_mtp_draft_model: + if self.is_mtp_draft_model: special_model_input["mtp_draft_input_hiddens"] = torch.randn( token_num, self.config["hidden_size"], dtype=self.data_type, device="cuda" ) @@ -1400,10 +1301,3 @@ def _gen_special_model_input(self, token_num: int): special_model_input["mtp_draft_input_hiddens"] = None return special_model_input - - def _gen_mtp_draft_special_model_input(self, token_num: int): - return { - "mtp_draft_input_hiddens": torch.randn( - token_num, self.config["hidden_size"], dtype=self.data_type, device="cuda" - ) - } diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 66ce4d0196..abf0e6cdd1 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -1,12 +1,7 @@ import torch -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Optional from typing import List -from lightllm.utils.envs_utils import ( - enable_diverse_mode_gqa_decode_fast_kernel, - enable_dynamic_mtp_verify, - enable_triton_mtp_kernel, -) from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor @@ -59,13 +54,8 @@ class ModelInput: # mtp_draft_input_hiddens 用于模型 mtp 模式下 # 的 draft 模型的输入 mtp_draft_input_hiddens: Optional[torch.Tensor] = None - # 部分 spec draft 模型会在服务 MTP 模式下执行普通 decode batch - # (例如 Eagle3 commit accepted rows),此时 attention 不应按 MTP 展开布局建参。 - disable_mtp_decode_att: bool = False - # Dynamic verification normally needs arbitrary per-request row groups. - # A full-width plan has the original fixed K+1 layout and can reuse the - # substantially cheaper Static MTP attention parameter construction. - use_static_mtp_layout: bool = False + # Maximum number of extra query rows per request. + draft_step: int = 0 def to_cuda(self): self.check_input() @@ -91,22 +81,10 @@ def to_cuda(self): if self.b_prefill_start_loc is not None: self.b_prefill_start_loc = self.b_prefill_start_loc.cuda(non_blocking=True) - if not self.is_prefill and enable_diverse_mode_gqa_decode_fast_kernel(): - batch_size = len(self.b_req_idx) - if self.b_mark_shared_group is None: - self.b_mark_shared_group = torch.ones(size=(batch_size,), dtype=torch.int32, device="cuda") - else: - self.b_mark_shared_group = self.b_mark_shared_group.cuda(non_blocking=True) - if self.b_shared_seq_len is None: - self.b_shared_seq_len = torch.zeros(size=(batch_size,), dtype=torch.int32, device="cuda") - else: - self.b_shared_seq_len = self.b_shared_seq_len.cuda(non_blocking=True) - elif not self.is_prefill and (enable_dynamic_mtp_verify() or enable_triton_mtp_kernel()): - batch_size = len(self.b_req_idx) - if self.b_mark_shared_group is None: - self.b_mark_shared_group = torch.ones(size=(batch_size,), dtype=torch.int32, device="cuda") - else: - self.b_mark_shared_group = self.b_mark_shared_group.cuda(non_blocking=True) + if self.b_mark_shared_group is not None: + self.b_mark_shared_group = self.b_mark_shared_group.cuda(non_blocking=True) + if self.b_shared_seq_len is not None: + self.b_shared_seq_len = self.b_shared_seq_len.cuda(non_blocking=True) def __post_init__(self): self.check_input() @@ -123,14 +101,10 @@ def check_input(self): class ModelOutput: # 通用变量 logits: torch.Tensor + # Hidden states collected for the active speculative strategy. + spec_hidden: Optional[torch.Tensor] = None # 用于判断 mem_indexes 是否成功写入 req manager 中的事件对象。 prefill_mem_indexes_ready_event: torch.Event = None - # DSpark dynamic verify 使用的 raw confidence logits,由 draft model - # post layer 产生,proposer 只负责 scatter 到 verify batch。 - mtp_draft_confidence_logits: Optional[torch.Tensor] = None - # TP-sharded DSpark Markov head 已经完成全局 greedy 选择时,直接携带 - # draft token,避免为了下游 argmax 再 all-gather 完整词表 logits。 - mtp_draft_token_ids: Optional[torch.Tensor] = None # prompt_logics 用于在开启 return_all_prompt_logics 模式(如 enable_prompt_logprobs)时, # 保存整个 prefill 阶段每一个 token 位置对应的 logits(而非仅最后一个位置的 logits)。 @@ -140,7 +114,5 @@ class ModelOutput: def to_no_ref_tensor(self): self.logits = tensor_to_no_ref_tensor(self.logits) - if self.mtp_draft_confidence_logits is not None: - self.mtp_draft_confidence_logits = tensor_to_no_ref_tensor(self.mtp_draft_confidence_logits) - if self.mtp_draft_token_ids is not None: - self.mtp_draft_token_ids = tensor_to_no_ref_tensor(self.mtp_draft_token_ids) + if self.spec_hidden is not None: + self.spec_hidden = tensor_to_no_ref_tensor(self.spec_hidden) diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 46eb8fdf4d..215caffdca 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -1,5 +1,4 @@ -import gc - +import os import torch import torch.distributed as dist import copy @@ -7,14 +6,12 @@ import triton from typing import Optional from lightllm.utils.log_utils import init_logger -from lightllm.utils.envs_utils import ( - get_env_start_args, - enable_dynamic_mtp_verify, - get_diverse_max_batch_shared_group_size, -) +from lightllm.utils.envs_utils import get_env_start_args +from lightllm.distributed import dist_group_manager from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.utils.torch_memory_saver_utils import ( TorchMemorySaverWrapper, + MemoryTag, ) from .infer_struct import InferStateInfo @@ -22,172 +19,68 @@ logger = init_logger(__name__) -def _build_mtp_mark_shared_group_values( - *, - batch_size: int, - mtp_group_size: int, - max_group_size: int, - split_groups: bool, -) -> list: - assert mtp_group_size > 0 - assert max_group_size > 0 +class CudaGraph: + # CudaGraph forward pass for the decoding stage. - group_cap = max_group_size if split_groups else mtp_group_size - b_mark_shared_group = [0 for _ in range(batch_size)] - for req_start in range(0, batch_size, mtp_group_size): - req_end = min(req_start + mtp_group_size, batch_size) - for group_start in range(req_start, req_end, group_cap): - group_size = min(group_cap, req_end - group_start) - b_mark_shared_group[group_start + group_size - 1] = group_size - return b_mark_shared_group + @staticmethod + def gen_cuda_graph_batch_sizes( + max_batch_size: int = 8, + tp_world_size: int = 1, + batch_multiplier: int = 1, + ): + args = get_env_start_args() + # gen cuda graph batch_sizes + # cuda graph gen for batch size = [1, 2, 3, ..., graph_split_batch_size] + # and [graph_split_batch_size + graph_grow_step_size, + # if the batch_multiplier is not 1, then the batch_sizes will be multiply of batch_multiplier -class CudaGraph: - # CudaGraph forward pass for the decoding stage. + split_size = args.graph_split_batch_size * batch_multiplier + grow_size = args.graph_grow_step_size * batch_multiplier + batch_sizes = [i * batch_multiplier for i in range(1, args.graph_split_batch_size + 1)] + batch_sizes.extend(range(split_size + grow_size, max_batch_size, grow_size)) + batch_sizes = sorted({size for size in batch_sizes if size < max_batch_size} | {max_batch_size}) + + if args.enable_tpsp_mix_mode: + batch_sizes = sorted({triton.cdiv(size, tp_world_size) * tp_world_size for size in batch_sizes}) + assert batch_sizes[-1] == max_batch_size + return batch_sizes def __init__( self, max_batch_size=8, max_len_in_batch=8192, tp_world_size: int = 1, - graph_split_batch_size: Optional[int] = None, - graph_grow_step_size: Optional[int] = None, + batch_multiplier: int = 1, + capture_infer_cost: bool = False, ): self.graph = {} self.tp_world_size = tp_world_size self.mempool = torch.cuda.graph_pool_handle() if torch.cuda.is_available() else None self.args = get_env_start_args() - self.raw_max_batch_size = max_batch_size - self.max_batch_size = None + self.max_batch_size = max_batch_size self.graph_max_len_in_batch = max_len_in_batch - self.graph_split_batch_size = int( - self.args.graph_split_batch_size if graph_split_batch_size is None else graph_split_batch_size - ) - self.graph_grow_step_size = int( - self.args.graph_grow_step_size if graph_grow_step_size is None else graph_grow_step_size - ) - assert self.graph_split_batch_size > 0 - assert self.graph_grow_step_size > 0 self.enable_decode_microbatch_overlap = self.args.enable_decode_microbatch_overlap self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) - self.spec_adapter = None - self.model = None - - self._refresh_cuda_graph_batch_sizes() - return - - def set_spec_adapter(self, spec_adapter, model=None): - # Decode graphs captured before the speculative runtime is attached - # use different batch sizes and cache keys, so none of them can be - # replayed afterwards. Release their private graph pools before the - # speculative warmup; retaining both generations can consume several - # extra GiB and OOM otherwise valid mem_fraction configurations. - if self.graph: - torch.cuda.synchronize() - self.graph.clear() - self.mempool = None - gc.collect() - torch.cuda.empty_cache() - self.mempool = torch.cuda.graph_pool_handle() - self.spec_adapter = spec_adapter - self.model = model - self._refresh_cuda_graph_batch_sizes() - return + self.capture_infer_cost = capture_infer_cost + self.infer_cost_ms_by_batch_size = {} - def _refresh_cuda_graph_batch_sizes(self): - # gen cuda graph batch_sizes - # cuda graph gen for batch size = [1, 2, 3, ..., graph_split_batch_size] - # and [graph_split_batch_size + graph_grow_step_size, - # if the mtp_step is not 0, then the batch_sizes will be multiply of (mtp_step + 1) - - mtp_step = self._get_decode_graph_mtp_step() - group_size = mtp_step + 1 - self.max_batch_size = (self.raw_max_batch_size // group_size) * group_size - assert self.max_batch_size > 0, "cuda graph max_batch_size must cover at least one decode group" - - graph_split_batch_size = self.graph_split_batch_size * group_size - graph_grow_step_size = self.graph_grow_step_size * group_size - - batch_sizes = [i * group_size for i in range(1, self.graph_split_batch_size + 1)] - for _batch_size in range( - graph_split_batch_size + graph_grow_step_size, - self.max_batch_size, - graph_grow_step_size, - ): - batch_sizes.append(_batch_size) - - batch_sizes = list(set([e for e in batch_sizes if e < self.max_batch_size])) - batch_sizes.append(self.max_batch_size) - batch_sizes.sort() - if self.args.enable_tpsp_mix_mode: - batch_sizes = [triton.cdiv(e, self.tp_world_size) * self.tp_world_size for e in batch_sizes] - batch_sizes = list(set(batch_sizes)) - batch_sizes.sort() - - self.cuda_graph_batch_sizes = batch_sizes - assert batch_sizes[-1] == self.max_batch_size + self.cuda_graph_batch_sizes = self.gen_cuda_graph_batch_sizes( + max_batch_size=self.max_batch_size, + tp_world_size=self.tp_world_size, + batch_multiplier=batch_multiplier, + ) logger.info(f"cuda graph batch_sizes: {self.cuda_graph_batch_sizes}") - return - - def _get_decode_graph_mtp_step(self) -> int: - if self.spec_adapter is None: - return self.args.mtp_step - return self.spec_adapter.get_decode_graph_mtp_step(self.model) - - def _get_decode_graph_warmup_mtp_step(self) -> int: - if self.spec_adapter is None: - return self.args.mtp_step - get_warmup_step = getattr(self.spec_adapter, "get_decode_graph_warmup_mtp_step", None) - if get_warmup_step is None: - return self.spec_adapter.get_decode_graph_mtp_step(self.model) - return get_warmup_step(self.model) - - def _is_block_draft_model(self) -> bool: - if self.spec_adapter is None or self.model is None: - return False - is_block_draft_model = getattr(self.spec_adapter, "is_block_draft_model", None) - return is_block_draft_model is not None and is_block_draft_model(self.model) def can_run(self, batch_size, max_len_in_batch): return batch_size <= self.max_batch_size and max_len_in_batch <= self.graph_max_len_in_batch - def _graph_key( - self, - batch_size: int, - model_context=None, - model_context1=None, - ): - spec_key = None - spec_key1 = None - if self.spec_adapter is not None and model_context is not None: - spec_key = self.spec_adapter.graph_cache_key(model_context, model=self.model) - if self.spec_adapter is not None and model_context1 is not None: - spec_key1 = self.spec_adapter.graph_cache_key(model_context1, model=self.model) - if spec_key is None and spec_key1 is None: - return batch_size - return (batch_size, spec_key, spec_key1) - - def _export_spec_capture(self, infer_state: InferStateInfo): - if self.spec_adapter is None: - return None - return self.spec_adapter.export_graph_capture() - - def _restore_spec_capture(self, infer_state: InferStateInfo, captured_hiddens) -> None: - if self.spec_adapter is not None: - self.spec_adapter.restore_graph_capture(captured_hiddens) - return - - def need_capture( - self, - batch_size, - model_context: Optional[ModelInput] = None, - model_context1: Optional[ModelInput] = None, - ): + def need_capture(self, batch_size): find_batch_size = self.find_closest_graph_batch_size(batch_size) - return ( - find_batch_size is not None - and self._graph_key(find_batch_size, model_context, model_context1) not in self.graph - ) + if find_batch_size is not None: + return find_batch_size not in self.graph + else: + assert False, "dead code" def find_closest_graph_batch_size(self, batch_size): index = bisect.bisect_left(self.cuda_graph_batch_sizes, batch_size) @@ -197,34 +90,6 @@ def find_closest_graph_batch_size(self, batch_size): else: return None - def _make_warmup_mtp_index(self, batch_size: int) -> torch.Tensor: - mtp_step = self._get_decode_graph_warmup_mtp_step() - if mtp_step <= 0: - return torch.zeros(batch_size, dtype=torch.int32, device="cuda") - return torch.arange(batch_size, dtype=torch.int32, device="cuda") % (mtp_step + 1) - - def _make_warmup_seq_len(self, batch_size: int) -> torch.Tensor: - mtp_step = self._get_decode_graph_warmup_mtp_step() - if mtp_step > 0: - group_size = mtp_step + 1 - return torch.arange(batch_size, dtype=torch.int32, device="cuda") % group_size + 2 - return torch.full((batch_size,), 2, dtype=torch.int32, device="cuda") - - def _make_warmup_mtp_mark_shared_group(self, batch_size: int) -> Optional[torch.Tensor]: - mtp_step = self._get_decode_graph_warmup_mtp_step() - if mtp_step <= 0: - return None - - mtp_group_size = mtp_step + 1 - max_group_size = get_diverse_max_batch_shared_group_size() - b_mark_shared_group = _build_mtp_mark_shared_group_values( - batch_size=batch_size, - mtp_group_size=mtp_group_size, - max_group_size=max_group_size, - split_groups=not self._is_block_draft_model(), - ) - return torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cuda") - def _capture_decode(self, decode_func, infer_state: InferStateInfo): graph_obj = torch.cuda.CUDAGraph() input_ids = infer_state.input_ids @@ -252,58 +117,11 @@ def _capture_decode(self, decode_func, infer_state: InferStateInfo): with self.torch_memory_saver.cuda_graph(graph_obj, pool=self.mempool): model_output = decode_func(infer_state) - spec_capture = self._export_spec_capture(infer_state) - self.graph[self._graph_key(batch_size, infer_state)] = (graph_obj, infer_state, model_output, spec_capture) + self.graph[batch_size] = (graph_obj, infer_state, model_output) graph_obj.replay() - self._record_capture_replay_infer_cost_ms( - graph_obj=graph_obj, - batch_size=batch_size, - is_draft_model=self._is_draft_model_capture(infer_state), - ) + self._measure_replay_cost(graph_obj=graph_obj, batch_size=batch_size) return model_output - def _record_capture_replay_infer_cost_ms( - self, - graph_obj: torch.cuda.CUDAGraph, - batch_size: int, - is_draft_model: bool, - ) -> None: - if not enable_dynamic_mtp_verify(): - return - if is_draft_model and self.args.mtp_mode == "dspark": - # DSpark's planner uses target verify cost plus confidence-derived - # capacity estimates. Draft block cost is not part of the decision, - # so avoid adding a runtime barrier/synchronize on lazy draft graph - # capture. - return - - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - dist.barrier(group=dist.group.WORLD) - start_event = torch.cuda.Event(enable_timing=True) - end_event = torch.cuda.Event(enable_timing=True) - graph_obj.replay() - start_event.record() - graph_obj.replay() - end_event.record() - end_event.synchronize() - infer_cost_ms_tensor = torch.tensor([start_event.elapsed_time(end_event)], dtype=torch.float32, device="cuda") - dist.all_reduce(infer_cost_ms_tensor, op=dist.ReduceOp.MIN, group=dist.group.WORLD) - infer_cost_ms = float(infer_cost_ms_tensor.item()) - g_infer_context.record_dynamic_mtp_infer_cost( - batch_size=batch_size, - infer_cost_ms=infer_cost_ms, - is_draft_model=is_draft_model, - ) - return - - def _is_draft_model_capture(self, infer_state: InferStateInfo) -> bool: - if infer_state.mtp_draft_input_hiddens is not None or getattr(infer_state, "is_draft_model", False): - return True - if self.spec_adapter is None or self.model is None: - return False - return self.spec_adapter.is_draft_model(self.model) - def _capture_decode_overlap( self, decode_func, @@ -334,25 +152,37 @@ def _capture_decode_overlap( with self.torch_memory_saver.cuda_graph(graph_obj, pool=self.mempool): model_output, model_output1 = decode_func(infer_state, infer_state1) - spec_capture = self._export_spec_capture(infer_state) - spec_capture1 = self._export_spec_capture(infer_state1) - self.graph[self._graph_key(batch_size, infer_state, infer_state1)] = ( + self.graph[batch_size] = ( graph_obj, infer_state, infer_state1, model_output, model_output1, - spec_capture, - spec_capture1, ) graph_obj.replay() - self._record_capture_replay_infer_cost_ms( - graph_obj=graph_obj, - batch_size=batch_size, - is_draft_model=self._is_draft_model_capture(infer_state), - ) + self._measure_replay_cost(graph_obj=graph_obj, batch_size=batch_size) return model_output, model_output1 + def _measure_replay_cost(self, graph_obj: torch.cuda.CUDAGraph, batch_size: int) -> None: + if not self.capture_infer_cost: + return + + dist.barrier(group=dist.group.WORLD) + start_event = torch.cuda.Event(enable_timing=True) + end_event = torch.cuda.Event(enable_timing=True) + graph_obj.replay() + start_event.record() + graph_obj.replay() + end_event.record() + end_event.synchronize() + infer_cost_ms_tensor = torch.tensor( + [start_event.elapsed_time(end_event)], + dtype=torch.float32, + device="cuda", + ) + dist.all_reduce(infer_cost_ms_tensor, op=dist.ReduceOp.MIN, group=dist.group.WORLD) + self.infer_cost_ms_by_batch_size[batch_size] = float(infer_cost_ms_tensor.item()) + def capture_decode( self, decode_func, @@ -371,10 +201,9 @@ def capture_decode( def _replay(self, infer_state: InferStateInfo): batch_size = infer_state.input_ids.shape[0] - graph_obj, graph_infer_state, graph_output, spec_capture = self.graph[self._graph_key(batch_size, infer_state)] + graph_obj, graph_infer_state, graph_output = self.graph[batch_size] graph_infer_state.copy_for_cuda_graph(infer_state) graph_obj.replay() - self._restore_spec_capture(infer_state, spec_capture) return graph_output def _replay_overlap( @@ -389,14 +218,10 @@ def _replay_overlap( graph_infer_state1, graph_model_output, graph_model_output1, - spec_capture, - spec_capture1, - ) = self.graph[self._graph_key(batch_size, infer_state, infer_state1)] + ) = self.graph[batch_size] graph_infer_state.copy_for_cuda_graph(infer_state) graph_infer_state1.copy_for_cuda_graph(infer_state1) graph_obj.replay() - self._restore_spec_capture(infer_state, spec_capture) - self._restore_spec_capture(infer_state1, spec_capture1) return graph_model_output, graph_model_output1 def replay(self, infer_state, infer_state1=None): @@ -413,18 +238,21 @@ def warmup(self, model): from .basemodel import TpPartBaseModel model: TpPartBaseModel = model + draft_step = model.decode_batch_multiplier - 1 + # decode cuda graph init for batch_size in self.cuda_graph_batch_sizes[::-1]: + seq_len = 2 + total_token_num = batch_size * seq_len max_len_in_batch = self.graph_max_len_in_batch input_ids = torch.tensor([1 for _ in range(batch_size)], dtype=torch.int64, device="cuda") mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) - b_seq_len = self._make_warmup_seq_len(batch_size) - total_token_num = int(b_seq_len.sum().item()) - b_mtp_index = self._make_warmup_mtp_index(batch_size) - b_mark_shared_group = self._make_warmup_mtp_mark_shared_group(batch_size) + b_seq_len = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda") + b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_mark_shared_group = torch.zeros(batch_size, dtype=torch.int32, device="cuda") model_input = ModelInput( batch_size=batch_size, @@ -438,27 +266,13 @@ def warmup(self, model): b_mtp_index=b_mtp_index, b_mark_shared_group=b_mark_shared_group, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), + draft_step=draft_step, is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], **model._gen_special_model_input(batch_size), ) model_output: ModelOutput = model.forward(model_input) del model_output - if ( - enable_dynamic_mtp_verify() - and self.args.mtp_mode in {"eagle3", "dspark"} - and self.spec_adapter is not None - and not self.spec_adapter.is_draft_model(model) - and batch_size % (self.args.mtp_step + 1) == 0 - ): - # Dynamic planners can switch back to the fixed K+1 target - # layout for a full-width iteration. Capture that graph - # variant on the dummy warmup state. Lazy capture on a live - # request would execute the in-place linear-attention state - # updates multiple times and corrupt subsequent decode state. - model_input.use_static_mtp_layout = True - model_output = model.forward(model_input) - del model_output del input_ids del mem_indexes del b_req_idx @@ -484,20 +298,23 @@ def warmup_overlap(self, model): from .basemodel import TpPartBaseModel model: TpPartBaseModel = model + draft_step = model.decode_batch_multiplier - 1 + for batch_size in self.cuda_graph_batch_sizes[::-1]: decode_batches = [] for micro_batch_index in [0, 1]: # dummy decoding, capture the cudagraph + seq_len = 2 + total_token_num = batch_size * seq_len max_len_in_batch = self.graph_max_len_in_batch input_ids = torch.tensor([1 for _ in range(batch_size)], dtype=torch.int64, device="cuda") mem_indexes = model.mem_manager.alloc(len(input_ids)).cuda() b_req_idx = torch.tensor( [model.req_manager.HOLD_REQUEST_ID for _ in range(batch_size)], dtype=torch.int32, device="cuda" ) - b_seq_len = self._make_warmup_seq_len(batch_size) - total_token_num = int(b_seq_len.sum().item()) - b_mtp_index = self._make_warmup_mtp_index(batch_size) - b_mark_shared_group = self._make_warmup_mtp_mark_shared_group(batch_size) + b_seq_len = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda") + b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_mark_shared_group = torch.zeros(batch_size, dtype=torch.int32, device="cuda") micro_batch = ModelInput( is_prefill=False, @@ -512,6 +329,7 @@ def warmup_overlap(self, model): b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), + draft_step=draft_step, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], **model._gen_special_model_input(batch_size), ) diff --git a/lightllm/common/basemodel/hidden_collector.py b/lightllm/common/basemodel/hidden_collector.py new file mode 100644 index 0000000000..a03fa90e59 --- /dev/null +++ b/lightllm/common/basemodel/hidden_collector.py @@ -0,0 +1,130 @@ +from __future__ import annotations + +from typing import Iterable, List, Optional + +import torch + + +def unpad_collected_hidden(hidden: Optional[torch.Tensor], token_count: int) -> Optional[torch.Tensor]: + return None if hidden is None else hidden[:token_count] + + +class NoopHiddenCollector: + """Null object used by models that do not expose speculative features.""" + + def add(self, layer_index: int, hidden: torch.Tensor) -> None: + return + + def prefill_outputs(self, final_hidden: torch.Tensor) -> List[torch.Tensor]: + return [final_hidden] + + def finish( + self, + infer_state, + final_hidden: torch.Tensor, + forward_outputs: Optional[List[torch.Tensor]] = None, + ) -> Optional[torch.Tensor]: + return None + + +class FinalHiddenCollector(NoopHiddenCollector): + """Returns the final decoder hidden state without per-layer bookkeeping.""" + + def finish( + self, + infer_state, + final_hidden: torch.Tensor, + forward_outputs: Optional[List[torch.Tensor]] = None, + ) -> torch.Tensor: + return final_hidden.contiguous() + + +class LayerHiddenCollector(NoopHiddenCollector): + """Collects selected decoder-layer outputs for an intermediate-hidden draft.""" + + def __init__(self, model, layer_ids: Iterable[int]) -> None: + self.model = model + self.layer_num = model.layers_num + self.layer_ids = frozenset(int(layer_id) for layer_id in layer_ids) + assert self.layer_ids, "layer hidden collector requires at least one layer id" + self.layer_hiddens: List[torch.Tensor] = [] + + def add(self, layer_index: int, hidden: torch.Tensor) -> None: + if layer_index not in self.layer_ids: + return + # Most LightLLM layers reuse their input buffer. Preserve intermediate + # layers while allowing the final layer output to remain zero-copy. + self.layer_hiddens.append(hidden if layer_index == self.layer_num - 1 else hidden.clone()) + + def _local_hidden(self) -> torch.Tensor: + assert len(self.layer_hiddens) == len( + self.layer_ids + ), f"captured {len(self.layer_hiddens)} hidden layers, expected {len(self.layer_ids)}" + if len(self.layer_hiddens) == 1: + return self.layer_hiddens[0] + return torch.cat(self.layer_hiddens, dim=-1) + + def prefill_outputs(self, final_hidden: torch.Tensor) -> List[torch.Tensor]: + return [final_hidden, self._local_hidden()] + + def finish( + self, + infer_state, + final_hidden: torch.Tensor, + forward_outputs: Optional[List[torch.Tensor]] = None, + ) -> torch.Tensor: + local_hidden = self._local_hidden() if forward_outputs is None else forward_outputs[1] + self.layer_hiddens.clear() + hidden = self.model.pre_infer._tpsp_allgather(input=local_hidden, infer_state=infer_state) + if infer_state.need_dp_prefill_balance: + hidden = infer_state._all_to_all_unbalance_get(data=hidden) + return hidden.contiguous() + + +class HiddenCollector: + """Collect hidden states for one or more independently executed microbatches.""" + + def __init__( + self, + model=None, + spec_mode: Optional[str] = None, + layer_ids: Optional[Iterable[int]] = None, + microbatch_count: int = 1, + ) -> None: + assert microbatch_count > 0 + layer_ids = None if layer_ids is None else tuple(layer_ids) + + collector_type = NoopHiddenCollector + collector_kwargs = {} + if spec_mode is not None: + assert model is not None + if model.is_mtp_draft_model: + if spec_mode not in ("dspark", "dflash"): + collector_type = FinalHiddenCollector + elif spec_mode in ("eagle3", "dspark", "dflash"): + assert layer_ids is not None + collector_type = LayerHiddenCollector + collector_kwargs = {"model": model, "layer_ids": layer_ids} + else: + collector_type = FinalHiddenCollector + + self.collectors = tuple(collector_type(**collector_kwargs) for _ in range(microbatch_count)) + + def add(self, layer_index: int, hidden: torch.Tensor, microbatch_index: int = 0) -> None: + self.collectors[microbatch_index].add(layer_index=layer_index, hidden=hidden) + + def finish( + self, + infer_state, + final_hidden: torch.Tensor, + forward_outputs: Optional[List[torch.Tensor]] = None, + microbatch_index: int = 0, + ) -> Optional[torch.Tensor]: + return self.collectors[microbatch_index].finish( + infer_state=infer_state, + final_hidden=final_hidden, + forward_outputs=forward_outputs, + ) + + def prefill_outputs(self, final_hidden: torch.Tensor, microbatch_index: int = 0) -> List[torch.Tensor]: + return self.collectors[microbatch_index].prefill_outputs(final_hidden) diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index 0ce3c2e0b9..ac9922ae69 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -1,17 +1,18 @@ import torch +import triton import collections from lightllm.common.kv_cache_mem_manager import MemoryManager +from lightllm.common.req_manager import ReqManager from lightllm.distributed import CustomProcessGroup -from typing import TYPE_CHECKING, Optional, List +from typing import Tuple, Any, Optional, List from .triton_kernel.gen_prefill_params import gen_prefill_params from .triton_kernel.gen_decode_params import gen_decode_params +from .triton_kernel.multimodal_emb import mark_multimodal_obj +from .batch_objs import ModelInput from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.dist_utils import get_global_dp_rank, get_dp_world_size from .attention import BasePrefillAttState, BaseDecodeAttState -if TYPE_CHECKING: - from lightllm.common.req_manager import ReqManager - class InferStateInfo: """ @@ -47,9 +48,12 @@ def __init__(self): # 的sum值, 其值等于 sum(b_ready_cache_len) self.prefix_total_token_num: int = None self.is_prefill: bool = None + # Specialized models override attention causality explicitly. + self.prefill_causal: bool = True + self.decode_causal: bool = True self.mem_manager: MemoryManager = None - self.req_manager: "ReqManager" = None + self.req_manager: ReqManager = None self.mem_index: torch.Tensor = None @@ -92,7 +96,8 @@ def __init__(self): # 在开启 mtp_mode 时,mtp draft model # 的输入会用到,其他模型和场景都不会用到 self.mtp_draft_input_hiddens: Optional[torch.Tensor] = None - self.disable_mtp_decode_att: bool = False + # Maximum number of extra query rows per request. + self.draft_step: int = 0 # 在单节点多dp的运行模式下,在进行prefill的阶段,如果出现了dp之间数据不平衡的现象, # 可以将推理的数据,进行重新分配到各个dp,在做 att 之前,重新 all to all 到各自的 @@ -130,19 +135,6 @@ def init_some_extra_state(self, model): ) = gen_decode_params(self.b_seq_len) self.b_kv_start_loc = self.b1_cu_kv_seq_len[0:-1] - @staticmethod - def build_draft_query_position_ids( - *, - selected_seq_len: torch.Tensor, - b_position_delta: Optional[torch.Tensor], - draft_step: int, - ) -> torch.Tensor: - offsets = torch.arange(draft_step, dtype=torch.long, device=selected_seq_len.device) - position_ids = selected_seq_len.to(dtype=torch.long)[:, None] + offsets[None, :] - if b_position_delta is not None: - position_ids = position_ids + b_position_delta.to(dtype=torch.long)[:, None] - return position_ids - def init_att_state(self): if self.is_prefill: self.prefill_att_state.init_state() diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index a4ccda2ff2..1c1148a55d 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -1,4 +1,6 @@ +import os import torch +import copy import bisect import triton from typing import List, Tuple @@ -6,6 +8,7 @@ from lightllm.utils.log_utils import init_logger from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor +from lightllm.distributed import dist_group_manager from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from .infer_struct import InferStateInfo from .cuda_graph import CudaGraph @@ -58,7 +61,10 @@ def can_run(self, handle_token_num: int): def need_capture(self, handle_token_num: int): finded_handle_token_num = self.find_closest_graph_handle_token_num(handle_token_num=handle_token_num) - return finded_handle_token_num is not None and finded_handle_token_num not in self.graph + if finded_handle_token_num is not None: + return finded_handle_token_num not in self.graph + else: + assert False, "dead code" def find_closest_graph_handle_token_num(self, handle_token_num: int): index = bisect.bisect_left(self.graph_handle_token_nums, handle_token_num) diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py similarity index 51% rename from lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py rename to lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py index 64db53d4df..a54aad84ab 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py @@ -9,25 +9,24 @@ def _fwd_kernel_cumprod_probs( req_to_next_token_probs, req_to_next_token_probs_stride, b_req_idx, - mtp_step, + max_draft_step, BLOCK_SIZE: tl.constexpr, ): cur_index = tl.program_id(0) - cur_req_idx = tl.load(b_req_idx + cur_index * (mtp_step + 1)) + cur_req_idx = tl.load(b_req_idx + cur_index * (max_draft_step + 1)) base_ptr = req_to_next_token_probs + cur_req_idx * req_to_next_token_probs_stride tl.store(base_ptr, 1.0) offset = tl.arange(0, BLOCK_SIZE) - store_mask = offset < (mtp_step + 1) + store_mask = offset < (max_draft_step + 1) probs = tl.load(base_ptr + offset, mask=store_mask, other=0.0) # offset 0 是 target sample,本轮恒接受;只有 draft 条件接受概率需要 clamp。 probs = tl.where(offset == 0, 1.0, probs) # 对于 draft probs 中大于 0.99 的值,设置为 0.99,避免错误的值,照成后续的采样操作失败。 # 对于 draft probs 中小于 0.01 的值,设置为 0.01,避免错误的值,照成后续的采样操作失败。 - # 这样修改后,我们可以做到,对于一个req的所有mtp step 步的概率的cumprod是降序的 - # 从而在采样时,我们只需要进行排序,则选中的位置,必然满足先后关系,避免一些复杂的 - # 额外操作。 + # This makes each request's cumulative acceptance probabilities monotonic, + # so global top-k selection cannot pick a later draft row without its prefix. probs = tl.where((offset != 0) & (probs >= 0.99), 0.99, probs) probs = tl.where((offset != 0) & (probs <= 0.01), 0.01, probs) @@ -99,23 +98,23 @@ def argsort(x, ids, dim: tl.core.constexpr = None, descending: tl.core.constexpr @triton.jit -def _fwd_kernel_sample_dynamic_mtp_steps( +def _fwd_kernel_select_dynamic_spec_rows( req_to_next_token_probs, req_to_next_token_probs_stride, - select_run_reqs, + selected_row_mask, b_req_idx, - mtp_step, - verify_step, + max_draft_step, + pre_draft_step, req_num, dynamic_batch_size, BLOCK_SIZE: tl.constexpr, ): - all_num = req_num * (verify_step + 1) + all_num = req_num * (pre_draft_step + 1) offset = tl.arange(0, BLOCK_SIZE) mask = offset < all_num - req_offset = offset // (verify_step + 1) - next_token_offset = offset % (verify_step + 1) - original_offset = req_offset * (mtp_step + 1) + next_token_offset + req_offset = offset // (pre_draft_step + 1) + next_token_offset = offset % (pre_draft_step + 1) + original_offset = req_offset * (max_draft_step + 1) + next_token_offset req_idx_index = tl.load(b_req_idx + original_offset, mask=mask, other=0) @@ -127,26 +126,26 @@ def _fwd_kernel_sample_dynamic_mtp_steps( sorted_probs, sorted_ids = argsort(probs, original_offset, descending=True) - tl.store(select_run_reqs + sorted_ids, 1, mask=mask & (offset < dynamic_batch_size)) + tl.store(selected_row_mask + sorted_ids, 1, mask=mask & (offset < dynamic_batch_size)) return -def sample_dynamic_mtp_req_mask( +def sample_dynamic_spec_row_mask( dynamic_batch_size: int, b_req_idx: torch.Tensor, req_to_next_token_probs: torch.Tensor, - mtp_step: int, - verify_step: int = None, + max_draft_step: int, + pre_draft_step: int = None, ) -> torch.Tensor: dynamic_batch_size = int(dynamic_batch_size) - mtp_step = int(mtp_step) - verify_step = mtp_step if verify_step is None else int(verify_step) - assert 0 <= verify_step <= mtp_step - assert b_req_idx.shape[0] % (mtp_step + 1) == 0 + max_draft_step = int(max_draft_step) + pre_draft_step = max_draft_step if pre_draft_step is None else int(pre_draft_step) + assert 0 <= pre_draft_step <= max_draft_step + assert b_req_idx.shape[0] % (max_draft_step + 1) == 0 assert req_to_next_token_probs.is_cuda assert dynamic_batch_size <= b_req_idx.shape[0] - req_num = len(b_req_idx) // (mtp_step + 1) - valid_row_num = req_num * (verify_step + 1) + req_num = len(b_req_idx) // (max_draft_step + 1) + valid_row_num = req_num * (pre_draft_step + 1) assert dynamic_batch_size <= valid_row_num # cumprod probs for each request @@ -154,27 +153,144 @@ def sample_dynamic_mtp_req_mask( req_to_next_token_probs=req_to_next_token_probs, req_to_next_token_probs_stride=req_to_next_token_probs.stride(0), b_req_idx=b_req_idx, - mtp_step=mtp_step, - BLOCK_SIZE=triton.next_power_of_2(mtp_step + 1), + max_draft_step=max_draft_step, + BLOCK_SIZE=triton.next_power_of_2(max_draft_step + 1), num_warps=1, num_stages=1, ) # 1 为选中, 0 为未选中 - select_run_reqs = torch.zeros((len(b_req_idx),), dtype=torch.int32, device="cuda") + selected_row_mask = torch.zeros((len(b_req_idx),), dtype=torch.int32, device="cuda") grid = (1,) - _fwd_kernel_sample_dynamic_mtp_steps[grid]( + _fwd_kernel_select_dynamic_spec_rows[grid]( req_to_next_token_probs=req_to_next_token_probs, req_to_next_token_probs_stride=req_to_next_token_probs.stride(0), - select_run_reqs=select_run_reqs, + selected_row_mask=selected_row_mask, b_req_idx=b_req_idx, - mtp_step=mtp_step, - verify_step=verify_step, + max_draft_step=max_draft_step, + pre_draft_step=pre_draft_step, req_num=req_num, dynamic_batch_size=dynamic_batch_size, BLOCK_SIZE=triton.next_power_of_2(valid_row_num), num_warps=1, num_stages=1, ) - return select_run_reqs + return selected_row_mask + + +@triton.jit +def _fwd_kernel_trim_post_sample_tensors( + b_req_idx, + out_b_req_idx, + b_temperatures, + out_b_temperatures, + b_top_ps, + out_b_top_ps, + b_top_ks, + out_b_top_ks, + b_length_penalty_param, + out_b_length_penalty_param, + b_mask_eos_reqs, + out_b_mask_eos_reqs, + selected_row_mask, + selected_dst_pos, + batch_size, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size + selected = tl.load(selected_row_mask + offsets, mask=mask, other=0) != 0 + dst_pos = tl.load(selected_dst_pos + offsets, mask=mask, other=0) + write_mask = mask & selected + + tl.store( + out_b_req_idx + dst_pos, + tl.load(b_req_idx + offsets, mask=mask, other=0), + mask=write_mask, + ) + tl.store( + out_b_temperatures + dst_pos, + tl.load(b_temperatures + offsets, mask=mask, other=0.0), + mask=write_mask, + ) + tl.store( + out_b_top_ps + dst_pos, + tl.load(b_top_ps + offsets, mask=mask, other=0.0), + mask=write_mask, + ) + tl.store( + out_b_top_ks + dst_pos, + tl.load(b_top_ks + offsets, mask=mask, other=0), + mask=write_mask, + ) + tl.store( + out_b_length_penalty_param + dst_pos, + tl.load(b_length_penalty_param + offsets, mask=mask, other=0), + mask=write_mask, + ) + tl.store( + out_b_mask_eos_reqs + dst_pos, + tl.load(b_mask_eos_reqs + offsets, mask=mask, other=0), + mask=write_mask, + ) + + +def trim_post_sample_tensors( + dynamic_batch_size: int, + selected_row_mask: torch.Tensor, + b_req_idx: torch.Tensor, + b_temperatures: torch.Tensor, + b_top_ps: torch.Tensor, + b_top_ks: torch.Tensor, + b_length_penalty_param: torch.Tensor, + b_mask_eos_reqs: torch.Tensor, +): + assert selected_row_mask.is_cuda + dynamic_batch_size = int(dynamic_batch_size) + selected_row_mask = selected_row_mask.to(torch.int32) + selected_dst_pos = torch.cumsum(selected_row_mask, dim=0, dtype=torch.int32) - 1 + batch_size = selected_row_mask.shape[0] + + output_tensors = tuple( + torch.empty((dynamic_batch_size,), dtype=tensor.dtype, device=tensor.device) + for tensor in ( + b_req_idx, + b_temperatures, + b_top_ps, + b_top_ks, + b_length_penalty_param, + b_mask_eos_reqs, + ) + ) + ( + out_b_req_idx, + out_b_temperatures, + out_b_top_ps, + out_b_top_ks, + out_b_length_penalty_param, + out_b_mask_eos_reqs, + ) = output_tensors + + block_size = 256 + _fwd_kernel_trim_post_sample_tensors[(triton.cdiv(batch_size, block_size),)]( + b_req_idx=b_req_idx, + out_b_req_idx=out_b_req_idx, + b_temperatures=b_temperatures, + out_b_temperatures=out_b_temperatures, + b_top_ps=b_top_ps, + out_b_top_ps=out_b_top_ps, + b_top_ks=b_top_ks, + out_b_top_ks=out_b_top_ks, + b_length_penalty_param=b_length_penalty_param, + out_b_length_penalty_param=out_b_length_penalty_param, + b_mask_eos_reqs=b_mask_eos_reqs, + out_b_mask_eos_reqs=out_b_mask_eos_reqs, + selected_row_mask=selected_row_mask, + selected_dst_pos=selected_dst_pos, + batch_size=batch_size, + BLOCK_SIZE=block_size, + num_warps=4, + num_stages=1, + ) + return output_tensors diff --git a/lightllm/common/basemodel/triton_kernel/fa3_utils.py b/lightllm/common/basemodel/triton_kernel/fa3_utils.py index c20d2443d5..154ca9637a 100644 --- a/lightllm/common/basemodel/triton_kernel/fa3_utils.py +++ b/lightllm/common/basemodel/triton_kernel/fa3_utils.py @@ -3,8 +3,8 @@ import triton.language as tl -_DYNAMIC_MTP_FA3_FAST_PATH_MAX_BATCH_SIZE = 1024 -_DYNAMIC_MTP_FA3_COMPACT_BLOCK_SIZE = 256 +_DYNAMIC_SPEC_FA3_FAST_PATH_MAX_BATCH_SIZE = 1024 +_DYNAMIC_SPEC_FA3_COMPACT_BLOCK_SIZE = 256 @triton.jit @@ -62,8 +62,41 @@ def page_table_copy( ) +def test_page_table_copy(): + import torch + + batch_size, seq_len = 2, 8 + + req_to_token_indexs = torch.arange(batch_size * seq_len, dtype=torch.int32).reshape(batch_size, seq_len).cuda() + + page_table = torch.full((batch_size, seq_len), -1, dtype=torch.int32, device="cuda") + + b_req_idx = torch.tensor([0, 2, 1, 3], dtype=torch.int32, device="cuda")[::2] + print(b_req_idx.stride()) + + page_table_copy(page_table, req_to_token_indexs, b_req_idx) + + print("req_to_token_indexs:") + print(req_to_token_indexs.cpu().numpy()) + print("b_req_idx:", b_req_idx.cpu().numpy()) + print("page_table:") + print(page_table.cpu().numpy()) + + for batch in range(batch_size): + src_idx = b_req_idx[batch].item() + expected = req_to_token_indexs[src_idx].cpu().numpy() + got = page_table[batch].cpu().numpy() + assert (expected == got).all(), f"Batch {batch} mismatch: expected {expected}, got {got}" + + print("✅ Test passed!") + + +if __name__ == "__main__": + test_page_table_copy() + + @triton.jit -def _build_dynamic_mtp_fa3_decode_params_kernel( +def _build_dynamic_spec_fa3_decode_params_kernel( b_req_idx, b_seq_len, b_mark_shared_group, @@ -79,6 +112,7 @@ def _build_dynamic_mtp_fa3_decode_params_kernel( mask = offsets < batch_size mark = tl.load(b_mark_shared_group + offsets, mask=mask, other=0) + # A positive mark closes a query group and stores the number of rows in it. is_group_end = mask & (mark > 0) dst_pos = tl.cumsum(tl.where(is_group_end, 1, 0), axis=0) - 1 @@ -97,7 +131,7 @@ def _build_dynamic_mtp_fa3_decode_params_kernel( @triton.jit -def _count_dynamic_mtp_fa3_decode_params_kernel( +def _count_dynamic_spec_fa3_decode_params_kernel( b_mark_shared_group, out_block_counts, batch_size, @@ -113,7 +147,7 @@ def _count_dynamic_mtp_fa3_decode_params_kernel( @triton.jit -def _compact_dynamic_mtp_fa3_decode_params_kernel( +def _compact_dynamic_spec_fa3_decode_params_kernel( b_req_idx, b_seq_len, b_mark_shared_group, @@ -132,6 +166,7 @@ def _compact_dynamic_mtp_fa3_decode_params_kernel( mask = offsets < batch_size mark = tl.load(b_mark_shared_group + offsets, mask=mask, other=0) + # A positive mark closes a query group and stores the number of rows in it. is_group_end = mask & (mark > 0) local_pos = tl.cumsum(tl.where(is_group_end, 1, 0), axis=0) - 1 block_count = tl.load(block_counts + block_id) @@ -149,25 +184,57 @@ def _compact_dynamic_mtp_fa3_decode_params_kernel( @torch.no_grad() -def build_dynamic_mtp_fa3_decode_params( +def build_dynamic_spec_fa3_decode_params( b_req_idx: torch.Tensor, b_seq_len: torch.Tensor, b_mark_shared_group: torch.Tensor, att_batch_size: int, hold_req_id: int, ): + """Convert dynamic speculative-verify rows to one metadata row per FA3 sequence. + + Query rows belonging to the same attention sequence are consecutive. + ``b_mark_shared_group`` is zero inside a group; its final row stores the + number of query rows in that group. + + For example, let ``H`` denote ``hold_req_id`` and suppose:: + + b_req_idx = [ 7, 7, 7, 4, 9, 9, H, H] + b_seq_len = [12, 13, 14, 8, 20, 21, 0, 0] + b_mark_shared_group = [ 0, 0, 3, 1, 0, 2, 0, 0] + + The positive marks close three attention sequences:: + + request 7 -> query length 3, final KV length 14 + request 4 -> query length 1, final KV length 8 + request 9 -> query length 2, final KV length 21 + + The operator retains each group's final row, compacts those rows to the + front, and pads the unused tail to ``att_batch_size``:: + + b_q_seq_len = [ 3, 1, 2, 0, 0, 0, 0, 0] + b_kv_seq_len = [14, 8, 21, 0, 0, 0, 0, 0] + b_att_req_idx = [7, 4, 9, H, H, H, H, H] + b_att_seq_len = [14, 8, 21, 0, 0, 0, 0, 0] + + Thus ``b_att_req_idx`` changes from one request id per query row to one id + per FA3 sequence; it selects that sequence's page-table row. A request may + still appear more than once if its rows are intentionally split into + multiple attention groups. Fixed output shapes let one CUDA Graph replay + different speculative verify layouts. + """ assert b_req_idx.is_cuda and b_seq_len.is_cuda and b_mark_shared_group.is_cuda assert b_req_idx.shape == b_seq_len.shape == b_mark_shared_group.shape assert b_req_idx.shape[0] == att_batch_size assert att_batch_size > 0 - if att_batch_size <= _DYNAMIC_MTP_FA3_FAST_PATH_MAX_BATCH_SIZE: + if att_batch_size <= _DYNAMIC_SPEC_FA3_FAST_PATH_MAX_BATCH_SIZE: b_q_seq_len = torch.empty((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) b_kv_seq_len = torch.empty((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) b_att_req_idx = torch.empty((att_batch_size,), dtype=torch.int32, device=b_req_idx.device) b_att_seq_len = torch.empty((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) - _build_dynamic_mtp_fa3_decode_params_kernel[(1,)]( + _build_dynamic_spec_fa3_decode_params_kernel[(1,)]( b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mark_shared_group=b_mark_shared_group, @@ -188,11 +255,11 @@ def build_dynamic_mtp_fa3_decode_params( b_att_req_idx = torch.full((att_batch_size,), hold_req_id, dtype=torch.int32, device=b_req_idx.device) b_att_seq_len = torch.zeros((att_batch_size,), dtype=torch.int32, device=b_seq_len.device) - block_size = _DYNAMIC_MTP_FA3_COMPACT_BLOCK_SIZE + block_size = _DYNAMIC_SPEC_FA3_COMPACT_BLOCK_SIZE grid = (triton.cdiv(att_batch_size, block_size),) block_counts = torch.empty((grid[0],), dtype=torch.int32, device=b_mark_shared_group.device) - _count_dynamic_mtp_fa3_decode_params_kernel[grid]( + _count_dynamic_spec_fa3_decode_params_kernel[grid]( b_mark_shared_group=b_mark_shared_group, out_block_counts=block_counts, batch_size=att_batch_size, @@ -202,7 +269,7 @@ def build_dynamic_mtp_fa3_decode_params( ) block_offsets = torch.cumsum(block_counts, dim=0, dtype=torch.int32) - _compact_dynamic_mtp_fa3_decode_params_kernel[grid]( + _compact_dynamic_spec_fa3_decode_params_kernel[grid]( b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mark_shared_group=b_mark_shared_group, diff --git a/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py b/lightllm/common/basemodel/triton_kernel/linear_att/spec_state_params.py similarity index 62% rename from lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py rename to lightllm/common/basemodel/triton_kernel/linear_att/spec_state_params.py index a9704310ec..623e1bfc08 100644 --- a/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py +++ b/lightllm/common/basemodel/triton_kernel/linear_att/spec_state_params.py @@ -4,7 +4,7 @@ @triton.jit -def _build_dynamic_mtp_linear_att_state_params_kernel( +def _build_dynamic_spec_linear_att_state_params_kernel( b_req_idx, b_mtp_index, req_to_mtp_state_index, @@ -22,15 +22,7 @@ def _build_dynamic_mtp_linear_att_state_params_kernel( req_idx = tl.load(b_req_idx + offsets, mask=token_mask, other=hold_req_id) mtp_index = tl.load(b_mtp_index + offsets, mask=token_mask, other=0) - # CUDA-graph warmup uses HOLD_REQUEST_ID for every row but still supplies - # a real 0..K MTP layout. Runtime graph padding, by contrast, appends HOLD - # rows whose mtp_index is always zero. Preserve warmup work while excluding - # runtime padding from the final real sequence. - has_hold_draft_row = tl.sum( - tl.where(token_mask & (req_idx == hold_req_id) & (mtp_index > 0), 1, 0), - axis=0, - ) > 0 - valid_row = token_mask & ((req_idx != hold_req_id) | has_hold_draft_row) + valid_row = token_mask & (req_idx != hold_req_id) actual_token_num = tl.sum(tl.where(valid_row, 1, 0), axis=0) # Keep tensor shapes dependent only on the target graph batch size. The @@ -41,6 +33,7 @@ def _build_dynamic_mtp_linear_att_state_params_kernel( tl.store(out_num_accepted_tokens + offsets, 1, mask=token_mask) tl.debug_barrier() + # mtp_index restarts at zero on the first row of every request group. is_start = valid_row & (mtp_index == 0) sequence_index = tl.cumsum(tl.where(is_start, 1, 0), axis=0) - 1 accepted_state_index = tl.load( @@ -54,18 +47,42 @@ def _build_dynamic_mtp_linear_att_state_params_kernel( tl.store(out_num_accepted_tokens + sequence_index, accepted_state_index + 1, mask=is_start) -def build_dynamic_mtp_linear_att_state_params( - *, +def build_dynamic_spec_linear_att_state_params( b_req_idx: torch.Tensor, b_mtp_index: torch.Tensor, req_to_mtp_state_index: torch.Tensor, hold_req_id: int, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Build graph-stable varlen GDN parameters for compact MTP verify rows. + """Convert compact speculative-verify rows to variable-length GDN sequences. - The compact input remains request-major and each request contributes a - prefix beginning at ``b_mtp_index == 0``. Outputs are padded to the token - batch size; trailing entries describe zero-length sequences. + Rows remain request-major after dynamic verification trimming, but each request can + contribute a different number of rows. ``b_mtp_index`` starts at zero for + every request and then increases within that request. + + For example, let ``H`` denote ``hold_req_id`` and suppose:: + + b_req_idx = [7, 7, 7, 4, 9, 9, H, H] + b_mtp_index = [0, 1, 2, 0, 0, 1, 0, 0] + + The first six rows contain three query sequences:: + + request 7 -> rows [0, 3), query length 3 + request 4 -> rows [3, 4), query length 1 + request 9 -> rows [4, 6), query length 2 + + Runtime ``H`` rows are graph padding and do not form query sequences. If + ``req_to_mtp_state_index`` stores ``{7: 2, 4: 0, 9: 1}``, the fixed-shape + outputs are:: + + b1_cu_q_seq_len = [0, 3, 4, 6, 6, 6, 6, 6, 6] + b_conv_buffer_idx = [7, 4, 9, H, H, H, H, H] + b_num_accepted_tokens = [3, 1, 2, 1, 1, 1, 1, 1] + + ``b_conv_buffer_idx`` therefore changes from one request id per query row + to one request id per GDN sequence. Repeated cumulative lengths describe + zero-length tail sequences. Keeping every output shape dependent only on + the padded input batch allows the same CUDA Graph to replay different + per-request query lengths. """ assert b_req_idx.is_cuda and b_mtp_index.is_cuda and req_to_mtp_state_index.is_cuda @@ -78,7 +95,7 @@ def build_dynamic_mtp_linear_att_state_params( b_conv_buffer_idx = torch.empty_like(b_req_idx) b_num_accepted_tokens = torch.empty_like(b_req_idx) - _build_dynamic_mtp_linear_att_state_params_kernel[(1,)]( + _build_dynamic_spec_linear_att_state_params_kernel[(1,)]( b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, req_to_mtp_state_index=req_to_mtp_state_index, diff --git a/lightllm/common/basemodel/triton_kernel/linear_att_copy.py b/lightllm/common/basemodel/triton_kernel/linear_att_copy.py index 0f2841ca2f..eb17507fb0 100644 --- a/lightllm/common/basemodel/triton_kernel/linear_att_copy.py +++ b/lightllm/common/basemodel/triton_kernel/linear_att_copy.py @@ -10,7 +10,6 @@ def _copy_linear_att_state_to_kv_buffer( cpu_kv_conv_ptr, # uint8 view: [buffer_num, linear_layer_num, conv_dim * cpu_conv_row_bytes] cpu_kv_ssm_ptr, # uint8 view: [buffer_num, linear_layer_num, ssm_bytes] b_req_idx, # [batch_size,] - req_to_mtp_state_index, # [max_request_num + 1,] big_page_buffer_ids, # [batch_size,] gpu_conv_stride_l, gpu_conv_stride_s, @@ -25,7 +24,7 @@ def _copy_linear_att_state_to_kv_buffer( cpu_kv_ssm_stride_s, cpu_kv_ssm_stride_l, cpu_kv_ssm_stride_d, - mtp_step: tl.constexpr, + mtp_step, gpu_conv_dim, # number of conv rows gpu_conv_tail_dim_bytes, # bytes copied per conv row; equals the CPU/cache row width gpu_ssm_tail_dim, @@ -51,10 +50,7 @@ def _copy_linear_att_state_to_kv_buffer( return cur_req_idx = tl.load(b_req_idx + cur_batch).to(tl.int64) - state_offset = 0 - if mtp_step > 0: - state_offset = tl.load(req_to_mtp_state_index + cur_req_idx).to(tl.int64) - cur_state_req_idx = (cur_req_idx * (mtp_step + 1) + state_offset).to(tl.int64) + cur_state_req_idx = (cur_req_idx * (mtp_step + 1)).to(tl.int64) gpu_conv_base = gpu_conv_ptr + cur_layer * gpu_conv_stride_l + cur_req_idx * gpu_conv_stride_s cpu_conv_base = cpu_kv_conv_ptr + big_page_buffer_idx * cpu_kv_conv_stride_s + cur_layer * cpu_kv_conv_stride_l @@ -84,7 +80,6 @@ def _copy_linear_att_state_to_kv_buffer( def copy_linear_att_state_to_kv_buffer( b_req_idx: torch.Tensor, - req_to_mtp_state_index: torch.Tensor, big_page_buffer_ids: torch.Tensor, gpu_conv_state: torch.Tensor, # [linear_layer_num, req_num, conv_dim, kernel_size] gpu_ssm_state: torch.Tensor, # [linear_layer_num, req_num * (mtp_step + 1), ...] @@ -94,10 +89,6 @@ def copy_linear_att_state_to_kv_buffer( ): # gpu_conv_state 的后两维可能是不连续的。 assert len(b_req_idx) == big_page_buffer_ids.shape[0] - if req_to_mtp_state_index is None: - assert mtp_step == 0 - # The constexpr branch below does not dereference this placeholder. - req_to_mtp_state_index = b_req_idx BLOCK = 4096 assert gpu_conv_state.dim() == 4, "gpu_conv_state must be [layer, s, conv_dim, widened_width]" @@ -138,7 +129,6 @@ def copy_linear_att_state_to_kv_buffer( cpu_kv_conv_ptr=cpu_kv_conv_state, cpu_kv_ssm_ptr=cpu_kv_ssm_state, b_req_idx=b_req_idx, - req_to_mtp_state_index=req_to_mtp_state_index, big_page_buffer_ids=big_page_buffer_ids, gpu_conv_stride_l=gpu_conv_state.stride(0), gpu_conv_stride_s=gpu_conv_state.stride(1), diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index c06337033e..f411db232c 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -4,8 +4,8 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput -from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import sample_dynamic_mtp_req_mask -from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size, get_env_start_args +from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import sample_dynamic_spec_row_mask +from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size @triton.jit @@ -13,7 +13,7 @@ def _fwd_kernel_mtp_verify( req_to_next_token_ids, req_to_next_token_ids_stride, new_next_token_ids, - mtp_accept_len, + spec_accept_len, b_req_mtp_start_loc, b_req_idx, accepted_index, @@ -42,7 +42,7 @@ def _fwd_kernel_mtp_verify( mismatch_positions = tl.where(match_mask, BLOCK_SIZE, offset) first_mismatch_pos = tl.min(mismatch_positions) accept_len = first_mismatch_pos + 1 - tl.store(mtp_accept_len + cur_index, accept_len) + tl.store(spec_accept_len + cur_index, accept_len) accpeted_index = tl.where((offset < accept_len), 1, 0) tl.store(accepted_index + req_offset, accpeted_index, mask=offset < req_mtp_num) return @@ -57,21 +57,21 @@ def mtp_verify( """ This function is used to verify the accept_len. Args: - req_to_next_token_ids: (max_req_num, max_mtp_step) + req_to_next_token_ids: (max_req_num, verify_width) b_req_mtp_start_loc: (num_reqs,) new_next_token_ids: (batch_size,) b_req_idx: (batch_size,) Returns: - mtp_accept_len: (num_reqs,) + spec_accept_len: (num_reqs,) accepted_index: (batch_size,) accepted_index: [1, 0, 1, 1, 0], 0 means the token is not accepted, 1 means the token is accepted. """ - max_mtp_step = req_to_next_token_ids.shape[1] + verify_width = req_to_next_token_ids.shape[1] BLOCK_SIZE = 16 - assert max_mtp_step <= BLOCK_SIZE, f"max_mtp_step must be less than {BLOCK_SIZE}" + assert verify_width <= BLOCK_SIZE, f"verify_width must be less than {BLOCK_SIZE}" num_reqs = b_req_mtp_start_loc.shape[0] req_mtp_all_num = b_req_idx.shape[0] - mtp_accept_len = torch.empty((num_reqs,), dtype=torch.int32, device=req_to_next_token_ids.device) + spec_accept_len = torch.empty((num_reqs,), dtype=torch.int32, device=req_to_next_token_ids.device) accepted_index = torch.empty((req_mtp_all_num,), dtype=torch.int32, device=req_to_next_token_ids.device) grid = (num_reqs,) @@ -80,7 +80,7 @@ def mtp_verify( req_to_next_token_ids=req_to_next_token_ids, req_to_next_token_ids_stride=req_to_next_token_ids.stride(0), new_next_token_ids=new_next_token_ids, - mtp_accept_len=mtp_accept_len, + spec_accept_len=spec_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, accepted_index=accepted_index, @@ -89,7 +89,7 @@ def mtp_verify( num_warps=num_warps, num_stages=1, ) - return mtp_accept_len, accepted_index + return spec_accept_len, accepted_index @triton.jit @@ -102,20 +102,21 @@ def _fwd_kernel_mtp_scatter_next_token_ids( req_to_next_token_probs_stride, all_next_token_probs, all_next_token_probs_stride, - mtp_accept_len, + spec_accept_len, b_req_mtp_start_loc, b_req_idx, mtp_step, - HAS_HAS_NEXT_TOKEN_PROBS: tl.constexpr, + HAS_NEXT_TOKEN_PROBS: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): + cur_index = tl.program_id(0) req_start_loc = tl.load(b_req_mtp_start_loc + cur_index) - accept_len = tl.load(mtp_accept_len + cur_index) + accept_len = tl.load(spec_accept_len + cur_index) cur_req_idx = tl.load(b_req_idx + req_start_loc) offset = tl.arange(0, BLOCK_SIZE) - if HAS_HAS_NEXT_TOKEN_PROBS: + if HAS_NEXT_TOKEN_PROBS: cur_next_token_probs = tl.load( all_next_token_probs + (req_start_loc + accept_len - 1) * all_next_token_probs_stride + offset, mask=offset < mtp_step, @@ -144,21 +145,21 @@ def mtp_scatter_next_token_ids( b_req_mtp_start_loc: torch.Tensor, all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, - mtp_accept_len: torch.Tensor, + spec_accept_len: torch.Tensor, req_to_next_token_probs: Optional[torch.Tensor] = None, all_next_token_probs: Optional[torch.Tensor] = None, ): - max_mtp_step = req_to_next_token_ids.shape[1] + verify_width = req_to_next_token_ids.shape[1] BLOCK_SIZE = 16 - assert max_mtp_step <= BLOCK_SIZE, f"max_mtp_step must be less than {BLOCK_SIZE}" + assert verify_width <= BLOCK_SIZE, f"verify_width must be less than {BLOCK_SIZE}" num_reqs = b_req_mtp_start_loc.shape[0] mtp_step = all_next_token_ids.shape[1] if req_to_next_token_probs is not None: assert all_next_token_probs is not None assert all_next_token_probs.shape == all_next_token_ids.shape - HAS_HAS_NEXT_TOKEN_PROBS = req_to_next_token_probs is not None - # Triton launch 参数阶段不能直接传 None;不开动态 MTP 时这里传一个不会被实际使用的 dummy tensor 即可。 + HAS_NEXT_TOKEN_PROBS = req_to_next_token_probs is not None + # Triton launch arguments cannot be None; static verification uses an unused placeholder. req_to_next_token_probs_arg = ( req_to_next_token_probs if req_to_next_token_probs is not None else req_to_next_token_ids ) @@ -181,11 +182,11 @@ def mtp_scatter_next_token_ids( req_to_next_token_probs_stride=req_to_next_token_probs_stride, all_next_token_probs=all_next_token_probs_arg, all_next_token_probs_stride=all_next_token_probs_stride, - mtp_accept_len=mtp_accept_len, + spec_accept_len=spec_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, mtp_step=mtp_step, - HAS_HAS_NEXT_TOKEN_PROBS=HAS_HAS_NEXT_TOKEN_PROBS, + HAS_NEXT_TOKEN_PROBS=HAS_NEXT_TOKEN_PROBS, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps, num_stages=1, @@ -193,7 +194,7 @@ def mtp_scatter_next_token_ids( @triton.jit -def _fwd_kernel_trim_dynamic_mtp_model_input( +def _fwd_kernel_compact_dynamic_spec_model_input( input_ids, out_input_ids, b_req_idx, @@ -312,35 +313,32 @@ def _fwd_kernel_pack_selected_rows_2d( tl.store(dst_ptrs, vals, mask=write_row & col_mask) -def _pack_selected_rows_2d( - src: Optional[torch.Tensor], - selected_mask_gpu: torch.Tensor, +def _pack_selected_hidden( + hidden: torch.Tensor, + selected_row_mask: torch.Tensor, selected_dst_pos: torch.Tensor, dynamic_batch_size: int, ): - if src is None: - return None - - assert src.is_cuda - assert src.ndim == 2 - assert selected_mask_gpu.is_cuda + assert hidden.is_cuda + assert hidden.ndim == 2 + assert selected_row_mask.is_cuda assert selected_dst_pos.is_cuda - assert src.shape[0] == selected_mask_gpu.shape[0] + assert hidden.shape[0] == selected_row_mask.shape[0] - selected_mask_gpu = selected_mask_gpu.to(torch.int32) - hidden_size = src.shape[1] - dst = torch.empty((dynamic_batch_size, hidden_size), dtype=src.dtype, device=src.device) - grid = (src.shape[0], triton.cdiv(hidden_size, 128)) + selected_row_mask = selected_row_mask.to(torch.int32) + hidden_size = hidden.shape[1] + dst = torch.empty((dynamic_batch_size, hidden_size), dtype=hidden.dtype, device=hidden.device) + grid = (hidden.shape[0], triton.cdiv(hidden_size, 128)) _fwd_kernel_pack_selected_rows_2d[grid]( - src=src, - src_stride_0=src.stride(0), - src_stride_1=src.stride(1), + src=hidden, + src_stride_0=hidden.stride(0), + src_stride_1=hidden.stride(1), dst=dst, dst_stride_0=dst.stride(0), dst_stride_1=dst.stride(1), - selected_mask=selected_mask_gpu, + selected_mask=selected_row_mask, selected_dst_pos=selected_dst_pos, - batch_size=src.shape[0], + batch_size=hidden.shape[0], hidden_size=hidden_size, BLOCK_N=128, num_warps=4, @@ -349,11 +347,12 @@ def _pack_selected_rows_2d( return dst -def _rebuild_trimmed_mtp_b_mark_shared_group_from_b_req_idx(b_req_idx: torch.Tensor) -> torch.Tensor: +def _rebuild_mtp_group_markers(b_req_idx: torch.Tensor, max_request_rows: int) -> torch.Tensor: assert b_req_idx.is_cuda batch_size = b_req_idx.shape[0] max_batch_shared_group_size = int(get_diverse_max_batch_shared_group_size()) assert max_batch_shared_group_size > 0 + assert max_request_rows > 0 if batch_size == 0: return torch.empty((0,), dtype=torch.int32, device=b_req_idx.device) @@ -364,7 +363,7 @@ def _rebuild_trimmed_mtp_b_mark_shared_group_from_b_req_idx(b_req_idx: torch.Ten out_b_mark_shared_group=b_mark_shared_group, batch_size=batch_size, max_batch_shared_group_size=max_batch_shared_group_size, - MAX_RUN_SCAN=16, + MAX_RUN_SCAN=max_request_rows - 1, BLOCK_SIZE=BLOCK_SIZE, num_warps=8, num_stages=1, @@ -372,19 +371,19 @@ def _rebuild_trimmed_mtp_b_mark_shared_group_from_b_req_idx(b_req_idx: torch.Ten return b_mark_shared_group -def _trim_decode_model_input_inplace( +def _compact_decode_model_input( model_input: ModelInput, - selected_mask_gpu: torch.Tensor, + selected_row_mask: torch.Tensor, dynamic_batch_size: int, ) -> ModelInput: assert not model_input.is_prefill - assert selected_mask_gpu.is_cuda + assert selected_row_mask.is_cuda assert model_input.b_req_idx.is_cuda assert model_input.b_mtp_index.is_cuda assert model_input.b_seq_len.is_cuda - # 动态 MTP 采样阶段已经保证 selected_mask_gpu 恰好选出 dynamic_batch_size 个位置。 - selected_mask_gpu = selected_mask_gpu.to(torch.int32) + # Dynamic scheduling guarantees exactly dynamic_batch_size selected rows. + selected_row_mask = selected_row_mask.to(torch.int32) old_batch_size = model_input.b_req_idx.shape[0] selected_dst_pos = torch.empty((old_batch_size,), dtype=torch.int32, device=model_input.b_req_idx.device) @@ -432,7 +431,7 @@ def _trim_decode_model_input_inplace( dummy_1d = model_input.b_req_idx BLOCK_SIZE = triton.next_power_of_2(old_batch_size) grid = (1,) - _fwd_kernel_trim_dynamic_mtp_model_input[grid]( + _fwd_kernel_compact_dynamic_spec_model_input[grid]( input_ids=model_input.input_ids if model_input.input_ids is not None else dummy_1d, out_input_ids=out_input_ids if out_input_ids is not None else dummy_1d, b_req_idx=model_input.b_req_idx, @@ -447,7 +446,7 @@ def _trim_decode_model_input_inplace( out_mem_indexes=out_mem_indexes if out_mem_indexes is not None else dummy_1d, b_shared_seq_len=model_input.b_shared_seq_len if model_input.b_shared_seq_len is not None else dummy_1d, out_b_shared_seq_len=out_b_shared_seq_len if out_b_shared_seq_len is not None else dummy_1d, - selected_mask=selected_mask_gpu, + selected_mask=selected_row_mask, selected_dst_pos=selected_dst_pos, batch_size=old_batch_size, HAS_INPUT_IDS=model_input.input_ids is not None, @@ -466,73 +465,60 @@ def _trim_decode_model_input_inplace( model_input.b_position_delta = out_b_position_delta model_input.mem_indexes = out_mem_indexes model_input.b_shared_seq_len = out_b_shared_seq_len - model_input.b_mark_shared_group = _rebuild_trimmed_mtp_b_mark_shared_group_from_b_req_idx(out_b_req_idx) + model_input.b_mark_shared_group = _rebuild_mtp_group_markers( + out_b_req_idx, + max_request_rows=model_input.draft_step + 1, + ) - if model_input.input_ids is not None: - assert model_input.input_ids.shape[0] == dynamic_batch_size if model_input.mtp_draft_input_hiddens is not None: assert model_input.mtp_draft_input_hiddens.is_cuda - model_input.mtp_draft_input_hiddens = _pack_selected_rows_2d( + model_input.mtp_draft_input_hiddens = _pack_selected_hidden( model_input.mtp_draft_input_hiddens, - selected_mask_gpu, + selected_row_mask, selected_dst_pos, dynamic_batch_size, ) - if model_input.b_position_delta is not None: - assert model_input.b_position_delta.shape[0] == dynamic_batch_size - if model_input.mem_indexes is not None: - assert model_input.mem_indexes.shape[0] == dynamic_batch_size model_input.batch_size = dynamic_batch_size return model_input -def prepare_dynamic_mtp_model_input( +def prepare_dynamic_spec_model_input( model_input: ModelInput, req_num: int, dynamic_batch_size: int, - req_to_next_token_ids: torch.Tensor, - req_to_next_token_probs: Optional[torch.Tensor] = None, - verify_step: Optional[int] = None, - use_prefix_selection: bool = False, + req_to_next_token_probs: torch.Tensor, + pre_draft_step: Optional[int] = None, ): - if req_to_next_token_probs is None: - selected_mask = torch.ones((model_input.batch_size,), dtype=torch.int32, device="cuda") - return model_input, selected_mask - req_num = int(req_num) dynamic_batch_size = int(dynamic_batch_size) - assert not model_input.is_prefill, "trim_dynamic_mtp_model_input only supports decode inputs" + assert not model_input.is_prefill, "prepare_dynamic_spec_model_input only supports decode inputs" + assert req_to_next_token_probs is not None assert dynamic_batch_size >= req_num assert dynamic_batch_size <= model_input.batch_size - mtp_step = int(get_env_start_args().mtp_step) - verify_step = mtp_step if verify_step is None else int(verify_step) - assert 0 <= verify_step <= mtp_step - assert model_input.batch_size == req_num * (mtp_step + 1) - assert dynamic_batch_size <= req_num * (verify_step + 1) - - # ! 在一个CUDA流上面的GPU操作会自动串行化,因此不需要额外同步 - # ! model_input必须在GPU上,才能高效进行trim操作 + max_draft_step = int(model_input.draft_step) + pre_draft_step = max_draft_step if pre_draft_step is None else int(pre_draft_step) + assert 0 <= pre_draft_step <= max_draft_step + assert model_input.batch_size == req_num * (max_draft_step + 1) + assert dynamic_batch_size <= req_num * (pre_draft_step + 1) + + # All compaction work stays on the current CUDA stream and needs no host sync. model_input.to_cuda() - if use_prefix_selection: - assert dynamic_batch_size == req_num * (verify_step + 1) - selected_mask_gpu = (model_input.b_mtp_index <= verify_step).to(dtype=torch.int32) - else: - selected_mask_gpu = sample_dynamic_mtp_req_mask( - dynamic_batch_size=dynamic_batch_size, - b_req_idx=model_input.b_req_idx, - req_to_next_token_probs=req_to_next_token_probs, - mtp_step=mtp_step, - verify_step=verify_step, - ) + selected_row_mask = sample_dynamic_spec_row_mask( + dynamic_batch_size=dynamic_batch_size, + b_req_idx=model_input.b_req_idx, + req_to_next_token_probs=req_to_next_token_probs, + max_draft_step=max_draft_step, + pre_draft_step=pre_draft_step, + ) - model_input = _trim_decode_model_input_inplace( + model_input = _compact_decode_model_input( model_input=model_input, - selected_mask_gpu=selected_mask_gpu, + selected_row_mask=selected_row_mask, dynamic_batch_size=dynamic_batch_size, ) - # Keep CPU mem_indexes unfiltered here. Copying selected_mask_gpu back to + # Keep CPU mem_indexes unfiltered here. Copying selected_row_mask back to # CPU in this hot path synchronizes the overlap stream; the router frees # unselected/rejected CPU mem indexes after its existing async mask copy is # consumed. Decode only needs b_position_delta on device, so placeholder @@ -543,18 +529,16 @@ def prepare_dynamic_mtp_model_input( empty_multimodal_params = {"images": [], "audios": []} model_input.multimodal_params = [empty_multimodal_params] * dynamic_batch_size - # ! 现在这些值没有实际作用,同时修改这些值还会导致阻塞操作,因此暂时不修改这些值了 - # model_input.total_token_num = int(model_input.b_seq_len.sum().item()) - # model_input.max_kv_seq_len = int(model_input.b_seq_len.max().item()) model_input.max_q_seq_len = 1 - return model_input, selected_mask_gpu + return model_input, selected_row_mask @triton.jit def _fwd_kernel_gen_b_req_mtp_start_loc( b_mtp_index, b_req_mtp_start_loc, - batch_size, + num_reqs: tl.constexpr, + batch_size: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): offset = tl.arange(0, BLOCK_SIZE) @@ -573,6 +557,7 @@ def gen_b_req_mtp_start_loc(b_mtp_index: torch.Tensor, num_reqs: int): _fwd_kernel_gen_b_req_mtp_start_loc[grid]( b_mtp_index=b_mtp_index, b_req_mtp_start_loc=b_req_mtp_start_loc, + num_reqs=num_reqs, batch_size=batch_size, BLOCK_SIZE=BLOCK_SIZE, num_warps=8, @@ -611,13 +596,13 @@ def _fwd_kernel_linear_att_mtp_state_index_update( return -def linear_att_mtp_state_index_update( +def linear_att_spec_state_index_update( req_to_mtp_state_index: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, b_req_idx: torch.Tensor, b_mtp_index: torch.Tensor, accepted_index: torch.Tensor, - max_mtp_step: int, + verify_width: int, ): """ Update req_to_mtp_state_index with the max b_mtp_index among accepted tokens per request. @@ -627,10 +612,10 @@ def linear_att_mtp_state_index_update( b_req_idx: (batch_size,) b_mtp_index: (batch_size,) accepted_index: (batch_size,), 1 means accepted, 0 means not accepted. - max_mtp_step: max mtp step per request, typically mtp_step + 1. + verify_width: maximum verify width per request, including the target row. """ BLOCK_SIZE = 16 - assert max_mtp_step <= BLOCK_SIZE, f"max_mtp_step must be less than {BLOCK_SIZE}" + assert verify_width <= BLOCK_SIZE, f"verify_width must be less than {BLOCK_SIZE}" num_reqs = b_req_mtp_start_loc.shape[0] req_mtp_all_num = b_req_idx.shape[0] @@ -649,3 +634,36 @@ def linear_att_mtp_state_index_update( num_warps=num_warps, num_stages=1, ) + + +def test_mtp_verify(): + req_to_next_token_ids = torch.tensor( + [[1, 2, -2, -1, -1], [1, 2, 0, -1, -1], [1, 3, 4, 4, 5]], dtype=torch.int64, device="cuda" + ) + b_req_idx = torch.tensor([0, 0, 2, 2, 2], dtype=torch.int32, device="cuda") + b_req_mtp_start_loc = torch.tensor([0, 2], dtype=torch.int32, device="cuda") + new_next_token_ids = torch.tensor([1, 4, 2, 4, 13], dtype=torch.int64, device="cuda") + all_next_token_ids = torch.tensor( + [[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], [13, 14, 15]], dtype=torch.int64, device="cuda" + ) + spec_accept_len, accepted_index = mtp_verify( + req_to_next_token_ids, b_req_mtp_start_loc, new_next_token_ids, b_req_idx + ) + mtp_scatter_next_token_ids( + req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, spec_accept_len + ) + print(spec_accept_len) + print(req_to_next_token_ids) + print(accepted_index) + + +def test_gen_b_req_mtp_start_loc(): + b_mtp_index = torch.tensor([0, 1, 0, 1, 2], dtype=torch.int32, device="cuda") + gt_output = torch.where(b_mtp_index == 0)[0] + b_req_mtp_start_loc = gen_b_req_mtp_start_loc(b_mtp_index, 2) + print(b_req_mtp_start_loc, gt_output) + + +if __name__ == "__main__": + test_mtp_verify() + # test_gen_b_req_mtp_start_loc() diff --git a/lightllm/common/basemodel/triton_kernel/norm/qk_norm.py b/lightllm/common/basemodel/triton_kernel/norm/qk_norm.py index e152a8dd83..e09d63f752 100644 --- a/lightllm/common/basemodel/triton_kernel/norm/qk_norm.py +++ b/lightllm/common/basemodel/triton_kernel/norm/qk_norm.py @@ -27,8 +27,8 @@ def _rms_norm_fwd_fused( var = tl.sum(x * x, axis=0) / head_dim rstd = 1 / tl.sqrt(var + eps) # Normalize and apply linear transformation - w = tl.load(W + tl.arange(0, BLOCK_SIZE)).to(tl.float32) - x_hat = x * rstd + w = tl.load(W + tl.arange(0, BLOCK_SIZE)) + x_hat = (x * rstd).to(X.dtype.element_ty) y = x_hat * w # Write output tl.store(X + cols, y.to(X.dtype.element_ty)) diff --git a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py index 68b3b09a46..49e2549265 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py +++ b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py @@ -190,8 +190,7 @@ def offload_gpu_kv_to_cpu_cache( return def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): - # Qwen3Next packs main full-attention layers first, followed by - # speculative draft layers. + # Qwen3Next 需要调整 layer_index layer_index = self.linear_config.get_full_att_kv_layer_index(layer_index) from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager diff --git a/lightllm/common/linear_att_cache_manager/config_objs.py b/lightllm/common/linear_att_cache_manager/config_objs.py index f588ec7d5c..6184437b9b 100644 --- a/lightllm/common/linear_att_cache_manager/config_objs.py +++ b/lightllm/common/linear_att_cache_manager/config_objs.py @@ -68,9 +68,9 @@ def get_conv_state_shape(self): # Base committed sliding-window state, without speculative MTP tail. return (self.get_conv_dim(), self.conv_kernel_size - 1) - def get_mtp_conv_state_shape(self, mtp_step: int): + def get_spec_conv_state_shape(self, max_draft_step: int): # Working state with room for S speculative tokens before acceptance. - return (self.get_conv_dim(), (self.conv_kernel_size - 1) + mtp_step) + return (self.get_conv_dim(), (self.conv_kernel_size - 1) + max_draft_step) def get_ssm_state_shape(self): return (self.num_linear_v_heads, self.head_linear_k_dim, self.head_linear_v_dim) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index f248cdbb39..41328277b1 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -7,7 +7,7 @@ from typing import List, Optional, TYPE_CHECKING from lightllm.common.basemodel.triton_kernel.gen_sampling_params import token_id_counter from lightllm.common.basemodel.triton_kernel.gen_sampling_params import update_req_to_token_id_counter -from lightllm.utils.envs_utils import get_env_start_args, enable_dynamic_mtp_verify +from lightllm.utils.envs_utils import get_env_start_args, enable_dynamic_spec from lightllm.utils.config_utils import get_vocab_size from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.common.linear_att_cache_manager.layer_cache import LayerCache @@ -116,20 +116,14 @@ def __init__(self, max_request_num): self.req_to_presence_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") self.req_to_frequency_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") self.req_to_repetition_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") - assert get_env_start_args().mtp_step <= 15, "mtp_step must be less than or equal to 15" self.req_to_next_token_ids = torch.zeros( - (max_request_num + 1, 16), + (max_request_num + 1, 8), dtype=torch.int64, device="cuda", ) - if enable_dynamic_mtp_verify(): - self.req_to_next_token_probs = torch.zeros( - (max_request_num + 1, 16), - dtype=torch.float32, - device="cuda", - ) - else: - self.req_to_next_token_probs = None + self.req_to_next_token_probs = ( + torch.zeros_like(self.req_to_next_token_ids, dtype=torch.float32) if enable_dynamic_spec() else None + ) self.req_to_exponential_decay_length_penalty = torch.zeros( max_request_num + 1, dtype=torch.float32, device="cuda" @@ -147,7 +141,7 @@ def __init__(self, max_request_num): def init_req_sampling_params(self, req: "InferReq"): shm_param = req.sampling_param.shm_param self.req_to_next_token_ids[req.req_idx][0:1].fill_(req.get_last_gen_token()) - if enable_dynamic_mtp_verify(): + if self.req_to_next_token_probs is not None: self.req_to_next_token_probs[req.req_idx].fill_(0.0) self.req_to_next_token_probs[req.req_idx][0:1].fill_(1.0) self.req_to_presence_penalty[req.req_idx].fill_(shm_param.presence_penalty) @@ -244,16 +238,16 @@ def gen_cpu_out_token_counter_sampling_params(self, req_objs: List["InferReq"]): class ReqManagerForMamba(ReqManager): def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_config: LinearAttCacheConfig): super().__init__(max_request_num, max_sequence_length, mem_manager) - self.mtp_step = get_env_start_args().mtp_step + self.max_draft_step = get_env_start_args().mtp_step # 因为在mtp的推理中,需要标记每个请求对应的mtp index状态(conv state 和 ssm state),在mtp对应序列中 # 的真实位置,所以需要需要一个标记来记录,不然算子无法找到真实的处理起点。 self.req_to_mtp_state_index = ( - torch.zeros((max_request_num + 1,), dtype=torch.int32, device="cuda") if self.mtp_step > 0 else None + torch.zeros((max_request_num + 1,), dtype=torch.int32, device="cuda") if self.max_draft_step > 0 else None ) # 突然想到, 在linear att 开启mtp的模式中,现在的prefill linear att 算子默认是从0的位置读取信息进行操作 # 所以不能支持 prefill decode mixed 操作了,因为一个decode过的请求,重新用prefill 算子跑,会出现读错linear # 状态位置的问题。导致bug, 在这里加个断言,以后可以支持上 TODO - if self.mtp_step > 0: + if self.max_draft_step > 0: assert get_env_start_args().enable_prefill_decode_mixed is False self.big_page_token_num = ( @@ -264,12 +258,12 @@ def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_con self.req_to_conv_state = LayerCache( size=(max_request_num + 1), dtype=self.linear_config.conv_state_dtype, - shape=self.linear_config.get_mtp_conv_state_shape(mtp_step=self.mtp_step), + shape=self.linear_config.get_spec_conv_state_shape(max_draft_step=self.max_draft_step), layer_num=self.linear_config.linear_layer_num, device="cuda", ) self.req_to_ssm_state = LayerCache( - size=(max_request_num + 1) * (self.mtp_step + 1), + size=(max_request_num + 1) * (self.max_draft_step + 1), dtype=self.linear_config.ssm_state_dtype, shape=self.linear_config.get_ssm_state_shape(), layer_num=self.linear_config.linear_layer_num, @@ -279,11 +273,11 @@ def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_con def init_linear_att_state(self, req: "InferReq"): conv_index = req.req_idx - ssm_start = req.req_idx * (self.mtp_step + 1) + ssm_start = req.req_idx * (self.max_draft_step + 1) self.req_to_conv_state.buffer[:, conv_index, ...].fill_(0) # #17: zero the FULL (mtp_step + 1)-row SSM block, not just canonical row +0, so a future # first-step verify reading offset>0 after fresh init never hits a never-written row (NaN). - self.req_to_ssm_state.buffer[:, ssm_start : ssm_start + (self.mtp_step + 1), ...].fill_(0) + self.req_to_ssm_state.buffer[:, ssm_start : ssm_start + (self.max_draft_step + 1), ...].fill_(0) if self.req_to_mtp_state_index is not None: self.req_to_mtp_state_index[req.req_idx] = 0 return @@ -304,7 +298,7 @@ def copy_big_page_buffer_to_linear_att_state(self, big_page_buffer_idx: int, req conv_state, ssm_state = big_page_buffers.get_state_cache(buffer_idx=big_page_buffer_idx) conv_dest = req.req_idx - ssm_dest = req.req_idx * (self.mtp_step + 1) + ssm_dest = req.req_idx * (self.max_draft_step + 1) conv_cache_width = conv_state.shape[-1] self.req_to_conv_state.buffer[:, conv_dest, ..., :conv_cache_width] = conv_state self.req_to_ssm_state.buffer[:, ssm_dest, ...] = ssm_state @@ -319,7 +313,7 @@ def copy_small_page_buffer_to_linear_att_state( buffer_idx=req.shared_kv_node.small_page_buffer_idx ) conv_dest = req.req_idx - ssm_dest = req.req_idx * (self.mtp_step + 1) + ssm_dest = req.req_idx * (self.max_draft_step + 1) conv_cache_width = conv_state.shape[-1] # TODO 下面这个从 cpu cache 拷贝数据的 gpu的操作,是否是阻塞的操作。 # 同时,非连续对象的拷贝,可能存在效率问题。 diff --git a/lightllm/common/speculative/__init__.py b/lightllm/common/speculative/__init__.py deleted file mode 100644 index 1388ca7cb3..0000000000 --- a/lightllm/common/speculative/__init__.py +++ /dev/null @@ -1,27 +0,0 @@ -from .config import ( - BlockDraftLayout, - SpeculativeConfig, - get_block_draft_layout, - normalize_speculative_draft_config, - is_dspark_draft_config, - is_eagle3_draft_config, - is_gemma4_dspark_draft_config, - is_qwen3_dflash_draft_config, - is_qwen3_5_dflash_draft_config, - is_qwen3_dspark_draft_config, - validate_dspark_family_draft_config, -) - -__all__ = [ - "BlockDraftLayout", - "SpeculativeConfig", - "get_block_draft_layout", - "normalize_speculative_draft_config", - "is_dspark_draft_config", - "is_eagle3_draft_config", - "is_gemma4_dspark_draft_config", - "is_qwen3_dflash_draft_config", - "is_qwen3_5_dflash_draft_config", - "is_qwen3_dspark_draft_config", - "validate_dspark_family_draft_config", -] diff --git a/lightllm/common/speculative/config.py b/lightllm/common/speculative/config.py deleted file mode 100644 index 3bd0581af0..0000000000 --- a/lightllm/common/speculative/config.py +++ /dev/null @@ -1,310 +0,0 @@ -from dataclasses import dataclass -from typing import Any, Mapping, MutableMapping, Optional - - -VANILLA_SPEC_MODES = frozenset({"vanilla_with_att", "vanilla_no_att", "qwen3next_vanilla"}) -EAGLE_SPEC_MODES = frozenset({"eagle_with_att", "eagle_no_att", "eagle3", "qwen3next_eagle"}) -BLOCK_SPEC_MODES = frozenset({"dspark", "dflash"}) -SPEC_MODES = VANILLA_SPEC_MODES | EAGLE_SPEC_MODES | BLOCK_SPEC_MODES - -ATTENTION_SPEC_MODES = frozenset({"vanilla_with_att", "eagle_with_att", "eagle3", "dspark", "dflash"}) -NO_ATTENTION_SPEC_MODES = frozenset({"vanilla_no_att", "eagle_no_att", "qwen3next_vanilla", "qwen3next_eagle"}) -TARGET_HIDDEN_SPEC_MODES = frozenset({"eagle3", "dspark", "dflash"}) -QWEN3_DFLASH_ARCHITECTURES = frozenset({"Qwen3DFlashModel", "Qwen3DSparkModel"}) -QWEN3_5_DFLASH_ARCHITECTURES = frozenset({"Qwen3_5DFlashModel"}) -QWEN3_DSPARK_ARCHITECTURES = frozenset({"Qwen3DSparkModel"}) -GEMMA4_DSPARK_ARCHITECTURES = frozenset({"Gemma4DSparkModel"}) -DSPARK_FAMILY_ARCHITECTURES = QWEN3_DFLASH_ARCHITECTURES | QWEN3_5_DFLASH_ARCHITECTURES | GEMMA4_DSPARK_ARCHITECTURES -DSPARK_MARKOV_HEAD_TYPES = frozenset({"vanilla", "gated", "rnn"}) -SPECULATIVE_CONFIG_SECTIONS = ("dflash_config", "dspark_config", "draft_config", "speculative_config", "mtp_config") -SPECULATIVE_DRAFT_CONFIG_KEYS = frozenset( - { - "block_size", - "target_layer_ids", - "mask_token_id", - "markov_rank", - "markov_head_type", - "enable_confidence_head", - "confidence_head_with_markov", - } -) - - -@dataclass(frozen=True) -class BlockDraftLayout: - """Runtime layout of a non-causal block draft checkpoint. - - ``query_block_size`` is the number of logits rows emitted for one anchor, - while ``proposal_output_start`` identifies the first row that represents a - draft token. Keeping both values explicit lets serving support checkpoints - whose query block includes a leading bonus row without teaching generic - scheduling or proposer code about a particular model architecture. - """ - - query_block_size: int - proposal_output_start: int - - @property - def draft_step(self) -> int: - return self.query_block_size - self.proposal_output_start - - def resolve_draft_step(self, configured_step: int) -> int: - """Use a positive configured step up to the checkpoint's proposal capacity.""" - configured_step = int(configured_step) - return configured_step if 0 < configured_step <= self.draft_step else self.draft_step - - -@dataclass(frozen=True) -class SpeculativeConfig: - """Normalized view of speculative decoding mode flags.""" - - mode: Optional[str] - step: int - dynamic_verify: bool = False - - @classmethod - def from_args(cls, args: Any, dynamic_verify: Optional[bool] = None) -> "SpeculativeConfig": - mode = getattr(args, "mtp_mode", None) - if dynamic_verify is None: - dynamic_verify = bool(getattr(args, "mtp_dynamic_verify", False)) - if mode == "dspark": - dynamic_verify = True - elif mode == "dflash": - dynamic_verify = False - return cls( - mode=mode, - step=int(getattr(args, "mtp_step", 0)), - dynamic_verify=dynamic_verify, - ) - - @property - def enabled(self) -> bool: - return self.mode is not None - - @property - def is_vanilla(self) -> bool: - return self.mode in VANILLA_SPEC_MODES - - @property - def is_eagle(self) -> bool: - return self.mode in EAGLE_SPEC_MODES - - @property - def is_eagle3(self) -> bool: - return self.mode == "eagle3" - - @property - def is_dspark(self) -> bool: - return self.mode == "dspark" - - @property - def is_dflash(self) -> bool: - return self.mode == "dflash" - - @property - def uses_block_draft_model(self) -> bool: - return self.mode in BLOCK_SPEC_MODES - - @property - def needs_target_layer_hidden(self) -> bool: - return self.mode in TARGET_HIDDEN_SPEC_MODES - - @property - def uses_attention_draft(self) -> bool: - return self.mode in ATTENTION_SPEC_MODES - - @property - def uses_no_attention_draft(self) -> bool: - return self.mode in NO_ATTENTION_SPEC_MODES - - @property - def uses_chained_draft_models(self) -> bool: - return self.mode in VANILLA_SPEC_MODES - - @property - def uses_recurrent_draft_model(self) -> bool: - return self.mode in EAGLE_SPEC_MODES - - @property - def draft_model_count(self) -> int: - if not self.enabled: - return 0 - return 1 if (self.uses_recurrent_draft_model or self.uses_block_draft_model) else self.step - - @property - def needs_draft_vocab_mapping(self) -> bool: - return self.is_eagle3 - - def get_decode_graph_mtp_step(self, *, model_config: Mapping[str, Any], is_draft_model: bool) -> int: - if (self.is_dflash or self.is_dspark) and is_draft_model: - return int(model_config["block_size"]) - 1 - if self.is_eagle3 and self.dynamic_verify: - # Dynamic Eagle3 physically compacts target rows and recurrent - # draft rows, so graph shapes must be available at unit batch - # granularity instead of only at multiples of the exposed depth. - return 0 - return self.step - - def get_decode_graph_warmup_mtp_step(self, *, model_config: Mapping[str, Any], is_draft_model: bool) -> int: - if self.is_eagle3 and self.dynamic_verify: - # Preserve representative MTP indices/shared-group metadata while - # retaining unit-granularity graph shape capture. - return min(3, self.step) - return self.get_decode_graph_mtp_step( - model_config=model_config, - is_draft_model=is_draft_model, - ) - - def validate(self) -> None: - if not self.enabled: - assert self.step == 0 - return - - assert self.mode in SPEC_MODES, f"unsupported speculative mode {self.mode}" - if not self.uses_block_draft_model: - assert self.step > 0 - else: - assert self.step >= 0 - if self.is_dspark: - assert self.dynamic_verify, "DSpark mode requires dynamic verify scheduling" - if self.uses_chained_draft_models: - assert self.draft_model_count == self.step - else: - assert self.draft_model_count == 1 - - -def is_eagle3_draft_config(config: Mapping[str, Any]) -> bool: - architectures = config.get("architectures", []) - return config.get("model_type") == "llama" or any( - architecture in ["Eagle3Speculator", "Qwen3Eagle3Model"] for architecture in architectures - ) - - -def is_dspark_draft_config(config: Mapping[str, Any]) -> bool: - architectures = config.get("architectures", []) - return any(architecture in DSPARK_FAMILY_ARCHITECTURES for architecture in architectures) - - -def is_qwen3_dflash_draft_config(config: Mapping[str, Any]) -> bool: - architectures = config.get("architectures", []) - return any(architecture in QWEN3_DFLASH_ARCHITECTURES for architecture in architectures) - - -def is_qwen3_5_dflash_draft_config(config: Mapping[str, Any]) -> bool: - architectures = config.get("architectures", []) - return any(architecture in QWEN3_5_DFLASH_ARCHITECTURES for architecture in architectures) - - -def is_qwen3_dspark_draft_config(config: Mapping[str, Any]) -> bool: - architectures = config.get("architectures", []) - return any(architecture in QWEN3_DSPARK_ARCHITECTURES for architecture in architectures) - - -def is_gemma4_dspark_draft_config(config: Mapping[str, Any]) -> bool: - architectures = config.get("architectures", []) - return any(architecture in GEMMA4_DSPARK_ARCHITECTURES for architecture in architectures) - - -def normalize_speculative_draft_config(config: MutableMapping[str, Any]) -> MutableMapping[str, Any]: - """Normalize supported speculative draft config layouts in place. - - LightLLM model code generally normalizes nested checkpoint config once at - load time, then downstream code reads a flat `network_config`. This helper - applies the same pattern to speculative draft checkpoints whose shared - fields may live under sections such as `dflash_config`. - """ - - for section in SPECULATIVE_CONFIG_SECTIONS: - nested_config = config.get(section) - if not isinstance(nested_config, Mapping): - continue - for key in SPECULATIVE_DRAFT_CONFIG_KEYS: - if key not in config and key in nested_config: - config[key] = nested_config[key] - return config - - -def validate_dspark_family_draft_config( - config: MutableMapping[str, Any], - *, - require_confidence_head: bool = False, -) -> None: - """Validate DFlash/DSpark checkpoint fields consumed by LightLLM serving.""" - - assert is_dspark_draft_config(config), f"unsupported DFlash/DSpark architecture: {config.get('architectures')}" - - normalize_speculative_draft_config(config) - - block_size = int(config.get("block_size", 0)) - assert block_size > 0, "DFlash/DSpark draft config must provide positive block_size" - - target_layer_ids = config.get("target_layer_ids") - assert ( - isinstance(target_layer_ids, (list, tuple)) and len(target_layer_ids) > 0 - ), "DFlash/DSpark draft config must provide non-empty target_layer_ids" - previous_layer_id = None - for raw_layer_id in target_layer_ids: - layer_id = int(raw_layer_id) - assert layer_id >= 0, ( - "LightLLM DFlash/DSpark serving expects decoder-layer target_layer_ids; " - "embedding-output layer_id=-1 is not supported" - ) - assert ( - previous_layer_id is None or layer_id > previous_layer_id - ), "DFlash/DSpark target_layer_ids must be strictly increasing" - previous_layer_id = layer_id - - assert "mask_token_id" in config, "DFlash/DSpark draft config must provide mask_token_id" - assert int(config["mask_token_id"]) >= 0, "DFlash/DSpark mask_token_id must be non-negative" - - markov_rank = int(config.get("markov_rank", 0)) - assert markov_rank >= 0, f"DFlash/DSpark markov_rank must be >= 0, got {markov_rank}" - if markov_rank > 0: - markov_head_type = str(config.get("markov_head_type", "")).lower() - assert ( - markov_head_type in DSPARK_MARKOV_HEAD_TYPES - ), f"unsupported DFlash/DSpark markov_head_type {markov_head_type!r}" - - enable_confidence_head = bool(config.get("enable_confidence_head", False)) - if require_confidence_head: - assert enable_confidence_head, "DSpark dynamic scheduling requires enable_confidence_head=true" - if enable_confidence_head: - assert ( - "confidence_head_with_markov" in config - ), "confidence_head_with_markov must be provided when enable_confidence_head is true" - if bool(config.get("confidence_head_with_markov", False)): - assert markov_rank > 0, "confidence_head_with_markov requires markov_rank > 0" - return - - -def get_block_draft_layout( - config: MutableMapping[str, Any], - *, - mode: str, - require_confidence_head: bool = False, -) -> BlockDraftLayout: - """Resolve the generic query/proposal layout of a block draft checkpoint. - - Proposal rows start at zero by default. Architectures with a different - upstream block contract are registered explicitly. - """ - - assert mode in BLOCK_SPEC_MODES, f"block draft layout is not defined for mode {mode!r}" - validate_dspark_family_draft_config( - config, - require_confidence_head=require_confidence_head, - ) - - query_block_size = int(config["block_size"]) - # The Z-Lab Qwen3.5 DFlash checkpoint defines block_size as the full - # query block: row 0 is the accepted/bonus query and proposals start at 1. - proposal_output_start = 1 if is_qwen3_5_dflash_draft_config(config) else 0 - - assert 0 <= proposal_output_start < query_block_size, ( - "block draft proposal_output_start must be within the query block: " - f"start={proposal_output_start}, block_size={query_block_size}" - ) - return BlockDraftLayout( - query_block_size=query_block_size, - proposal_output_start=proposal_output_start, - ) diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H800/_fwd_kernel_mtp_diverse_stage1_single_token:v1/{block_seq=256,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H800/_fwd_kernel_mtp_diverse_stage1_single_token:v1/{block_seq=256,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H800.json deleted file mode 100644 index 6e0ce74445..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.4.0/NVIDIA_H800/_fwd_kernel_mtp_diverse_stage1_single_token:v1/{block_seq=256,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H800.json +++ /dev/null @@ -1,326 +0,0 @@ -{ - "1000000032": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 4 - }, - "1000000064": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 8 - }, - "1000000128": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 8 - }, - "1000000256": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 8 - }, - "1000000512": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 8 - }, - "1000001024": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 8 - }, - "1000002048": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 8 - }, - "1000008192": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 8 - }, - "1000016384": { - "BLOCK_BATCH": 4, - "BLOCK_N": 32, - "num_stages": 3, - "num_warps": 4 - }, - "128000000032": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 2 - }, - "128000000064": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2 - }, - "128000000128": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 2 - }, - "128000000256": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 4, - "num_warps": 2 - }, - "128000000512": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 3, - "num_warps": 2 - }, - "128000001024": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 4 - }, - "128000002048": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 4 - }, - "128000008192": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 2 - }, - "128000016384": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 2 - }, - "16000000032": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 4 - }, - "16000000064": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 8 - }, - "16000000128": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 8 - }, - "16000000256": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 8 - }, - "16000000512": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 4 - }, - "16000001024": { - "BLOCK_BATCH": 4, - "BLOCK_N": 32, - "num_stages": 4, - "num_warps": 4 - }, - "16000002048": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 4, - "num_warps": 2 - }, - "16000008192": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 2 - }, - "16000016384": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 4 - }, - "32000000032": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 4 - }, - "32000000064": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 4 - }, - "32000000128": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 4 - }, - "32000000256": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 4 - }, - "32000000512": { - "BLOCK_BATCH": 4, - "BLOCK_N": 32, - "num_stages": 4, - "num_warps": 4 - }, - "32000001024": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 3, - "num_warps": 2 - }, - "32000002048": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 3, - "num_warps": 2 - }, - "32000008192": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 2 - }, - "32000016384": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 2 - }, - "64000000032": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 2 - }, - "64000000064": { - "BLOCK_BATCH": 4, - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 4 - }, - "64000000128": { - "BLOCK_BATCH": 4, - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 4 - }, - "64000000256": { - "BLOCK_BATCH": 4, - "BLOCK_N": 32, - "num_stages": 4, - "num_warps": 2 - }, - "64000000512": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 3, - "num_warps": 2 - }, - "64000001024": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 3, - "num_warps": 2 - }, - "64000002048": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 2 - }, - "64000008192": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 2 - }, - "64000016384": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 2 - }, - "8000000032": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 8 - }, - "8000000064": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 4 - }, - "8000000128": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 8 - }, - "8000000256": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 8 - }, - "8000000512": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 8 - }, - "8000001024": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 4, - "num_warps": 4 - }, - "8000002048": { - "BLOCK_BATCH": 4, - "BLOCK_N": 32, - "num_stages": 3, - "num_warps": 4 - }, - "8000008192": { - "BLOCK_BATCH": 4, - "BLOCK_N": 16, - "num_stages": 4, - "num_warps": 2 - }, - "8000016384": { - "BLOCK_BATCH": 4, - "BLOCK_N": 64, - "num_stages": 3, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index b730bccb8d..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=4,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,254 +0,0 @@ -{ - "1000000032": { - "BLOCK_N": 16, - "num_stages": 3, - "num_warps": 2, - "warp_specialize": true - }, - "1000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 4, - "warp_specialize": false - }, - "1000000128": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 8, - "warp_specialize": true - }, - "1000000256": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "1000000512": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 8, - "warp_specialize": false - }, - "1000001024": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "1000002048": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "128000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "128000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "128000000128": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "128000000256": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "128000000512": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "128000001024": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "128000002048": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "16000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "16000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "16000000128": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "16000000256": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "16000000512": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "16000001024": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "16000002048": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "32000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "32000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "32000000128": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "32000000256": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "32000000512": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "32000001024": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "32000002048": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "64000000128": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "64000000256": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000000512": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000001024": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000002048": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "8000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "8000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "8000000128": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 4, - "warp_specialize": false - }, - "8000000256": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "8000000512": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 8, - "warp_specialize": true - }, - "8000001024": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "8000002048": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=8,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=8,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index d04a6d8fe2..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage1_single_token:v2/{block_batch=4,gqa_group_size=8,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,254 +0,0 @@ -{ - "1000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "1000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 4, - "warp_specialize": false - }, - "1000000128": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "1000000256": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "1000000512": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 4, - "warp_specialize": false - }, - "1000001024": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "1000002048": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "128000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "128000000064": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "128000000128": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "128000000256": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "128000000512": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "128000001024": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "128000002048": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "16000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "16000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "16000000128": { - "BLOCK_N": 16, - "num_stages": 3, - "num_warps": 2, - "warp_specialize": false - }, - "16000000256": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "16000000512": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "16000001024": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "16000002048": { - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 4, - "warp_specialize": false - }, - "32000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "32000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "32000000128": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "32000000256": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "32000000512": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "32000001024": { - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 4, - "warp_specialize": false - }, - "32000002048": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000000064": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "64000000128": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "64000000256": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000000512": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "64000001024": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "64000002048": { - "BLOCK_N": 64, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "8000000032": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "8000000064": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "8000000128": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "8000000256": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": true - }, - "8000000512": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "8000001024": { - "BLOCK_N": 16, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - }, - "8000002048": { - "BLOCK_N": 32, - "num_stages": 2, - "num_warps": 2, - "warp_specialize": false - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index fcb4db67b1..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "1000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "128000000032": { - "num_stages": 1, - "num_warps": 2 - }, - "16000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "32000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "64000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "8000000128": { - "num_stages": 1, - "num_warps": 2 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index d5a8839f80..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=128,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "1000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "128000000032": { - "num_stages": 1, - "num_warps": 4 - }, - "16000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "32000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "64000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "8000000128": { - "num_stages": 1, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index d75cfb554e..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "1000000128": { - "num_stages": 1, - "num_warps": 2 - }, - "128000000032": { - "num_stages": 1, - "num_warps": 2 - }, - "16000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "32000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "64000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "8000000128": { - "num_stages": 1, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index a4b1a211cd..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=16,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "1000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "128000000032": { - "num_stages": 1, - "num_warps": 4 - }, - "16000000128": { - "num_stages": 1, - "num_warps": 2 - }, - "32000000064": { - "num_stages": 1, - "num_warps": 2 - }, - "64000000064": { - "num_stages": 1, - "num_warps": 2 - }, - "8000000128": { - "num_stages": 1, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index d5a8839f80..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "1000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "128000000032": { - "num_stages": 1, - "num_warps": 4 - }, - "16000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "32000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "64000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "8000000128": { - "num_stages": 1, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index 42975b4b9f..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=32,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "1000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "128000000032": { - "num_stages": 1, - "num_warps": 4 - }, - "16000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "32000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "64000000064": { - "num_stages": 1, - "num_warps": 2 - }, - "8000000128": { - "num_stages": 1, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index 0220ae94ed..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.bfloat16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "1000000128": { - "num_stages": 1, - "num_warps": 2 - }, - "128000000032": { - "num_stages": 1, - "num_warps": 2 - }, - "16000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "32000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "64000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "8000000128": { - "num_stages": 1, - "num_warps": 2 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json deleted file mode 100644 index 13615a6096..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.5.0/NVIDIA_H200/_fwd_kernel_mtp_diverse_stage2_single_token:v2/{block_n=64,out_dtype=torch.float16,q_head_dim=128}_NVIDIA_H200.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "1000000128": { - "num_stages": 1, - "num_warps": 4 - }, - "128000000032": { - "num_stages": 1, - "num_warps": 2 - }, - "16000000128": { - "num_stages": 1, - "num_warps": 2 - }, - "32000000064": { - "num_stages": 1, - "num_warps": 2 - }, - "64000000064": { - "num_stages": 1, - "num_warps": 4 - }, - "8000000128": { - "num_stages": 1, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json deleted file mode 100644 index f55b637832..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=16,H=12,K=128,V=128}_NVIDIA_H800.json +++ /dev/null @@ -1,8 +0,0 @@ -{ - "4": { - "BK": 128, - "BV": 128, - "num_stages": 4, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json deleted file mode 100644 index cc5c68eb79..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=32,H=12,K=128,V=128}_NVIDIA_H800.json +++ /dev/null @@ -1,8 +0,0 @@ -{ - "4": { - "BK": 128, - "BV": 64, - "num_stages": 2, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json deleted file mode 100644 index 7421097fa4..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_fwd_o/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json +++ /dev/null @@ -1,8 +0,0 @@ -{ - "4": { - "BK": 64, - "BV": 128, - "num_stages": 3, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_gated_delta_rule_fwd_h/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_gated_delta_rule_fwd_h/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json deleted file mode 100644 index d831f32c4a..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_gated_delta_rule_fwd_h/{BT=64,H=12,K=128,V=128}_NVIDIA_H800.json +++ /dev/null @@ -1,7 +0,0 @@ -{ - "4": { - "BV": 32, - "num_stages": 4, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_local_cumsum_scalar/{B=1,BT=64,H=12,IS_VARLEN=true,REVERSE=false}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_local_cumsum_scalar/{B=1,BT=64,H=12,IS_VARLEN=true,REVERSE=false}_NVIDIA_H800.json deleted file mode 100644 index 14509fffea..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_local_cumsum_scalar/{B=1,BT=64,H=12,IS_VARLEN=true,REVERSE=false}_NVIDIA_H800.json +++ /dev/null @@ -1,38 +0,0 @@ -{ - "1": { - "num_warps": 1 - }, - "1024": { - "num_warps": 2 - }, - "128": { - "num_warps": 2 - }, - "16": { - "num_warps": 8 - }, - "2048": { - "num_warps": 1 - }, - "256": { - "num_warps": 2 - }, - "32": { - "num_warps": 1 - }, - "4": { - "num_warps": 2 - }, - "4096": { - "num_warps": 2 - }, - "64": { - "num_warps": 1 - }, - "8": { - "num_warps": 1 - }, - "8192": { - "num_warps": 2 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_scaled_dot_kkt_fwd/{BT=64,H=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_scaled_dot_kkt_fwd/{BT=64,H=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json deleted file mode 100644 index a97cabf8b2..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/chunk_scaled_dot_kkt_fwd/{BT=64,H=12,IS_VARLEN=true,K=128}_NVIDIA_H800.json +++ /dev/null @@ -1,7 +0,0 @@ -{ - "4": { - "BK": 64, - "num_stages": 3, - "num_warps": 2 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/fused_gdn_gating:v1/{NUM_HEADS=12,a_dtype=torch.bfloat16}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/fused_gdn_gating:v1/{NUM_HEADS=12,a_dtype=torch.bfloat16}_NVIDIA_H800.json deleted file mode 100644 index 5284108618..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/fused_gdn_gating:v1/{NUM_HEADS=12,a_dtype=torch.bfloat16}_NVIDIA_H800.json +++ /dev/null @@ -1,50 +0,0 @@ -{ - "1": { - "BLK_HEADS": 64, - "num_warps": 2 - }, - "1024": { - "BLK_HEADS": 16, - "num_warps": 1 - }, - "128": { - "BLK_HEADS": 64, - "num_warps": 4 - }, - "16": { - "BLK_HEADS": 16, - "num_warps": 1 - }, - "2048": { - "BLK_HEADS": 64, - "num_warps": 2 - }, - "256": { - "BLK_HEADS": 16, - "num_warps": 1 - }, - "32": { - "BLK_HEADS": 8, - "num_warps": 2 - }, - "4": { - "BLK_HEADS": 4, - "num_warps": 1 - }, - "4096": { - "BLK_HEADS": 64, - "num_warps": 2 - }, - "64": { - "BLK_HEADS": 4, - "num_warps": 4 - }, - "8": { - "BLK_HEADS": 16, - "num_warps": 1 - }, - "8192": { - "BLK_HEADS": 16, - "num_warps": 1 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/gated_rmsnorm_forward:v1/{N=128,has_bias=false,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/gated_rmsnorm_forward:v1/{N=128,has_bias=false,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H800.json deleted file mode 100644 index 233215c4f2..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/gated_rmsnorm_forward:v1/{N=128,has_bias=false,weight_dtype=torch.bfloat16,x_dtype=torch.bfloat16}_NVIDIA_H800.json +++ /dev/null @@ -1,50 +0,0 @@ -{ - "12": { - "BLOCK_N": 128, - "num_warps": 1 - }, - "12288": { - "BLOCK_N": 128, - "num_warps": 1 - }, - "1536": { - "BLOCK_N": 512, - "num_warps": 1 - }, - "192": { - "BLOCK_N": 512, - "num_warps": 1 - }, - "24576": { - "BLOCK_N": 512, - "num_warps": 1 - }, - "3072": { - "BLOCK_N": 128, - "num_warps": 1 - }, - "384": { - "BLOCK_N": 64, - "num_warps": 2 - }, - "48": { - "BLOCK_N": 256, - "num_warps": 2 - }, - "49152": { - "BLOCK_N": 128, - "num_warps": 1 - }, - "768": { - "BLOCK_N": 64, - "num_warps": 2 - }, - "96": { - "BLOCK_N": 256, - "num_warps": 2 - }, - "98304": { - "BLOCK_N": 128, - "num_warps": 1 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/mrope_triton_fused:v1/{HEAD_DIM=256,K_HEAD_NUM=1,Q_HEAD_NUM=6,dtype=torch.bfloat16}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/mrope_triton_fused:v1/{HEAD_DIM=256,K_HEAD_NUM=1,Q_HEAD_NUM=6,dtype=torch.bfloat16}_NVIDIA_H800.json deleted file mode 100644 index 3e5f0d7165..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/mrope_triton_fused:v1/{HEAD_DIM=256,K_HEAD_NUM=1,Q_HEAD_NUM=6,dtype=torch.bfloat16}_NVIDIA_H800.json +++ /dev/null @@ -1,50 +0,0 @@ -{ - "1": { - "num_stages": 2, - "num_warps": 4 - }, - "1024": { - "num_stages": 4, - "num_warps": 4 - }, - "128": { - "num_stages": 2, - "num_warps": 1 - }, - "16": { - "num_stages": 4, - "num_warps": 4 - }, - "2048": { - "num_stages": 4, - "num_warps": 4 - }, - "256": { - "num_stages": 2, - "num_warps": 1 - }, - "32": { - "num_stages": 1, - "num_warps": 2 - }, - "4": { - "num_stages": 4, - "num_warps": 4 - }, - "4096": { - "num_stages": 2, - "num_warps": 2 - }, - "64": { - "num_stages": 2, - "num_warps": 2 - }, - "8": { - "num_stages": 4, - "num_warps": 4 - }, - "8192": { - "num_stages": 2, - "num_warps": 2 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/silu_and_mul_fwd:v1/{N=4352,out_dtype=torch.bfloat16}_NVIDIA_H800.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/silu_and_mul_fwd:v1/{N=4352,out_dtype=torch.bfloat16}_NVIDIA_H800.json deleted file mode 100644 index 6fc416c0ec..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H800/silu_and_mul_fwd:v1/{N=4352,out_dtype=torch.bfloat16}_NVIDIA_H800.json +++ /dev/null @@ -1,74 +0,0 @@ -{ - "1": { - "BLOCK_M": 128, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "1024": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "128": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "16": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2048": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "256": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "32": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 8 - }, - "4": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 4 - }, - "4096": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "64": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "8": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "8192": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - } -} \ No newline at end of file diff --git a/lightllm/models/__init__.py b/lightllm/models/__init__.py index bd8db7f8dd..e547068b1b 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -1,175 +1,110 @@ -from importlib import import_module - -from .registry import get_model as _registry_get_model -from .registry import get_model_class as _registry_get_model_class - - -_MODEL_EXPORTS = { - "MixtralTpPartModel": ("lightllm.models.mixtral.model", "MixtralTpPartModel"), - "BloomTpPartModel": ("lightllm.models.bloom.model", "BloomTpPartModel"), - "LlamaTpPartModel": ("lightllm.models.llama.model", "LlamaTpPartModel"), - "StarcoderTpPartModel": ("lightllm.models.starcoder.model", "StarcoderTpPartModel"), - "Starcoder2TpPartModel": ("lightllm.models.starcoder2.model", "Starcoder2TpPartModel"), - "QWenTpPartModel": ("lightllm.models.qwen.model", "QWenTpPartModel"), - "Qwen2TpPartModel": ("lightllm.models.qwen2.model", "Qwen2TpPartModel"), - "Qwen3TpPartModel": ("lightllm.models.qwen3.model", "Qwen3TpPartModel"), - "Qwen3MOEModel": ("lightllm.models.qwen3_moe.model", "Qwen3MOEModel"), - "Qwen3NextTpPartModel": ("lightllm.models.qwen3next.model", "Qwen3NextTpPartModel"), - "InternlmTpPartModel": ("lightllm.models.internlm.model", "InternlmTpPartModel"), - "StablelmTpPartModel": ("lightllm.models.stablelm.model", "StablelmTpPartModel"), - "Internlm2TpPartModel": ("lightllm.models.internlm2.model", "Internlm2TpPartModel"), - "Internlm2RewardTpPartModel": ( - "lightllm.models.internlm2_reward.model", - "Internlm2RewardTpPartModel", - ), - "MistralTpPartModel": ("lightllm.models.mistral.model", "MistralTpPartModel"), - "MiniCPMTpPartModel": ("lightllm.models.minicpm.model", "MiniCPMTpPartModel"), - "LlavaTpPartModel": ("lightllm.models.llava.model", "LlavaTpPartModel"), - "QWenVLTpPartModel": ("lightllm.models.qwen_vl.model", "QWenVLTpPartModel"), - "Gemma_2bTpPartModel": ("lightllm.models.gemma_2b.model", "Gemma_2bTpPartModel"), - "Phi3TpPartModel": ("lightllm.models.phi3.model", "Phi3TpPartModel"), - "Deepseek2TpPartModel": ("lightllm.models.deepseek2.model", "Deepseek2TpPartModel"), - "Deepseek3_2TpPartModel": ("lightllm.models.deepseek3_2.model", "Deepseek3_2TpPartModel"), - "Glm4MoeLiteTpPartModel": ( - "lightllm.models.glm4_moe_lite.model", - "Glm4MoeLiteTpPartModel", - ), - "InternVLLlamaTpPartModel": ("lightllm.models.internvl.model", "InternVLLlamaTpPartModel"), - "InternVLPhi3TpPartModel": ("lightllm.models.internvl.model", "InternVLPhi3TpPartModel"), - "InternVLQwen2TpPartModel": ("lightllm.models.internvl.model", "InternVLQwen2TpPartModel"), - "InternVLDeepSeek2TpPartModel": ( - "lightllm.models.internvl.model", - "InternVLDeepSeek2TpPartModel", - ), - "InternVLInternlm2TpPartModel": ( - "lightllm.models.internvl.model", - "InternVLInternlm2TpPartModel", - ), - "Qwen2VLTpPartModel": ("lightllm.models.qwen2_vl.model", "Qwen2VLTpPartModel"), - "Qwen2RewardTpPartModel": ("lightllm.models.qwen2_reward.model", "Qwen2RewardTpPartModel"), - "Qwen3VLTpPartModel": ("lightllm.models.qwen3_vl.model", "Qwen3VLTpPartModel"), - "Qwen3VLMOETpPartModel": ("lightllm.models.qwen3_vl_moe.model", "Qwen3VLMOETpPartModel"), - "Gemma3TpPartModel": ("lightllm.models.gemma3.model", "Gemma3TpPartModel"), - "Gemma4TpPartModel": ("lightllm.models.gemma4.model", "Gemma4TpPartModel"), - "Tarsier2Qwen2TpPartModel": ("lightllm.models.tarsier2.model", "Tarsier2Qwen2TpPartModel"), - "Tarsier2Qwen2VLTpPartModel": ( - "lightllm.models.tarsier2.model", - "Tarsier2Qwen2VLTpPartModel", - ), - "Tarsier2LlamaTpPartModel": ("lightllm.models.tarsier2.model", "Tarsier2LlamaTpPartModel"), - "GptOssTpPartModel": ("lightllm.models.gpt_oss.model", "GptOssTpPartModel"), - "Qwen3OmniMOETpPartModel": ( - "lightllm.models.qwen3_omni_moe_thinker.model", - "Qwen3OmniMOETpPartModel", - ), - "Qwen3_5TpPartModel": ("lightllm.models.qwen3_5.model", "Qwen3_5TpPartModel"), - "Qwen3_5MOETpPartModel": ("lightllm.models.qwen3_5_moe.model", "Qwen3_5MOETpPartModel"), - "Qwen3_5DFlashModel": ("lightllm.models.qwen3_5_dflash.model", "Qwen3_5DFlashModel"), - "Qwen3_5DSparkModel": ("lightllm.models.qwen3_5_dspark.model", "Qwen3_5DSparkModel"), +from lightllm.models.mixtral.model import MixtralTpPartModel +from lightllm.models.bloom.model import BloomTpPartModel +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.starcoder.model import StarcoderTpPartModel +from lightllm.models.starcoder2.model import Starcoder2TpPartModel +from lightllm.models.qwen.model import QWenTpPartModel +from lightllm.models.qwen2.model import Qwen2TpPartModel +from lightllm.models.qwen3.model import Qwen3TpPartModel +from lightllm.models.qwen3_moe.model import Qwen3MOEModel +from lightllm.models.qwen3next.model import Qwen3NextTpPartModel +from lightllm.models.internlm.model import InternlmTpPartModel +from lightllm.models.stablelm.model import StablelmTpPartModel +from lightllm.models.internlm2.model import Internlm2TpPartModel +from lightllm.models.internlm2_reward.model import Internlm2RewardTpPartModel +from lightllm.models.mistral.model import MistralTpPartModel +from lightllm.models.minicpm.model import MiniCPMTpPartModel +from lightllm.models.llava.model import LlavaTpPartModel +from lightllm.models.qwen_vl.model import QWenVLTpPartModel +from lightllm.models.gemma_2b.model import Gemma_2bTpPartModel +from lightllm.models.phi3.model import Phi3TpPartModel +from lightllm.models.deepseek2.model import Deepseek2TpPartModel +from lightllm.models.deepseek3_2.model import Deepseek3_2TpPartModel +from lightllm.models.glm4_moe_lite.model import Glm4MoeLiteTpPartModel +from lightllm.models.internvl.model import ( + InternVLLlamaTpPartModel, + InternVLPhi3TpPartModel, + InternVLQwen2TpPartModel, + InternVLDeepSeek2TpPartModel, +) +from lightllm.models.internvl.model import InternVLInternlm2TpPartModel +from lightllm.models.qwen2_vl.model import Qwen2VLTpPartModel +from lightllm.models.qwen2_reward.model import Qwen2RewardTpPartModel +from lightllm.models.qwen3_vl.model import Qwen3VLTpPartModel +from lightllm.models.qwen3_vl_moe.model import Qwen3VLMOETpPartModel +from lightllm.models.gemma3.model import Gemma3TpPartModel +from lightllm.models.gemma4.model import Gemma4TpPartModel +from lightllm.models.tarsier2.model import ( + Tarsier2Qwen2TpPartModel, + Tarsier2Qwen2VLTpPartModel, + Tarsier2LlamaTpPartModel, +) +from lightllm.models.gpt_oss.model import GptOssTpPartModel +from lightllm.models.qwen3_omni_moe_thinker.model import Qwen3OmniMOETpPartModel +from lightllm.models.qwen3_5.model import Qwen3_5TpPartModel +from lightllm.models.qwen3_5_moe.model import Qwen3_5MOETpPartModel +from lightllm.models.deepseek_mtp.model import Deepseek3MTPModel +from lightllm.models.glm4_moe_lite_mtp.model import Glm4MoeLiteMTPModel +from lightllm.models.mistral_mtp.model import MistralMTPModel +from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel +from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel +from lightllm.models.qwen3_5_moe_mtp.model import Qwen3_5MoeMTPModel +from lightllm.models.qwen3_5_mtp.model import Qwen3_5MTPModel +from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel +from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel +from lightllm.models.qwen3_eagle.model import Qwen3EagleModel +from lightllm.models.qwen3_moe_mtp.model import Qwen3MOEMTPModel +from .registry import get_model, get_model_class + + +_ATTENTION_DRAFT_MODELS = { + "deepseek_v3": Deepseek3MTPModel, + "glm4_moe_lite": Glm4MoeLiteMTPModel, + "qwen3_5": Qwen3_5MTPModel, + "qwen3_5_text": Qwen3_5MTPModel, + "qwen3_5_moe": Qwen3_5MoeMTPModel, + "qwen3_5_moe_text": Qwen3_5MoeMTPModel, } -_MODEL_TYPE_REGISTRY_MODULES = { - "starcoder2": ("lightllm.models.starcoder2.model",), - "internlm2": ("lightllm.models.internlm2.model",), - "llava": ("lightllm.models.llava.model",), - "qwen": ("lightllm.models.qwen.model",), - "qwen2": ("lightllm.models.qwen2.model",), - "qwen2_vl": ("lightllm.models.qwen2_vl.model",), - "qwen2_5_vl": ("lightllm.models.qwen2_vl.model",), - "qwen3": ("lightllm.models.qwen3.model",), - "qwen3_moe": ("lightllm.models.qwen3_moe.model",), - "qwen3_next": ("lightllm.models.qwen3next.model",), - "qwen3_vl": ("lightllm.models.qwen3_vl.model",), - "qwen3_vl_moe": ("lightllm.models.qwen3_vl_moe.model",), - "qwen3_omni_moe": ("lightllm.models.qwen3_omni_moe_thinker.model",), - "qwen3_5": ("lightllm.models.qwen3_5.model",), - "qwen3_5_moe": ("lightllm.models.qwen3_5_moe.model",), - "deepseek_v2": ("lightllm.models.deepseek2.model",), - "deepseek_v3": ("lightllm.models.deepseek2.model",), - "deepseek_v32": ("lightllm.models.deepseek3_2.model",), - "glm4_moe_lite": ("lightllm.models.glm4_moe_lite.model",), - "bloom": ("lightllm.models.bloom.model",), - "gpt_bigcode": ("lightllm.models.starcoder.model",), - "minicpm": ("lightllm.models.minicpm.model",), - "gemma3": ("lightllm.models.gemma3.model",), - "gemma": ("lightllm.models.gemma_2b.model",), - "gemma4": ("lightllm.models.gemma4.model",), - "internlm": ("lightllm.models.internlm.model",), - "stablelm": ("lightllm.models.stablelm.model",), - "mistral": ("lightllm.models.mistral.model",), - "gpt_oss": ("lightllm.models.gpt_oss.model",), - "phi3": ("lightllm.models.phi3.model",), - "llama": ("lightllm.models.llama.model",), - "mixtral": ("lightllm.models.mixtral.model",), - "internvl_chat": ("lightllm.models.internvl.model",), +_NO_ATTENTION_DRAFT_MODELS = { + "mistral": MistralMTPModel, + "qwen3_moe": Qwen3MOEMTPModel, } -_bootstrapped_registry_modules = set() - - -def _load_model_attr(name): - module_name, attr_name = _MODEL_EXPORTS[name] - value = getattr(import_module(module_name), attr_name) - globals()[name] = value - return value - -def _has_architecture(model_cfg: dict, name: str) -> bool: - return any(name in architecture for architecture in model_cfg.get("architectures", [])) - -def _llava_text_model_type(model_cfg: dict) -> str: - return model_cfg.get("llm_config", {}).get("model_type", "") or model_cfg.get("text_config", {}).get( - "model_type", "" +def get_draft_model_class(model_cfg, spec_mode, is_linear_att_mixed_model=None): + architectures = set(model_cfg.get("architectures", ())) + + if spec_mode == "eagle3" and "Qwen3Eagle3Model" in architectures: + return Qwen3EagleModel + + if spec_mode == "dflash": + if "Qwen3_5DFlashModel" in architectures: + if is_linear_att_mixed_model is False: + raise ValueError("Qwen3_5DFlashModel requires a linear-attention mixed target") + return Qwen3_5DFlashModel + if architectures.intersection(("Qwen3DFlashModel", "Qwen3DSparkModel")): + if is_linear_att_mixed_model is True: + raise ValueError("linear-attention mixed targets require a Qwen3_5DFlashModel checkpoint") + return Qwen3DFlashModel + + if spec_mode == "dspark" and "Qwen3DSparkModel" in architectures: + if is_linear_att_mixed_model: + return Qwen3_5DSparkModel + return Qwen3DSparkModel + + model_type = model_cfg.get("model_type", "") + if model_type in _ATTENTION_DRAFT_MODELS: + if spec_mode not in ("vanilla_with_att", "eagle_with_att", "eagle3", "dspark", "dflash"): + raise ValueError(f"{model_type} requires an attention draft mode, got {spec_mode}") + return _ATTENTION_DRAFT_MODELS[model_type] + + if model_type in _NO_ATTENTION_DRAFT_MODELS: + if spec_mode not in ("vanilla_no_att", "eagle_no_att", "qwen3next_vanilla", "qwen3next_eagle"): + raise ValueError(f"{model_type} requires a no-attention draft mode, got {spec_mode}") + return _NO_ATTENTION_DRAFT_MODELS[model_type] + + raise ValueError( + f"Unsupported speculative draft model: mode={spec_mode}, " + f"model_type={model_cfg.get('model_type')}, architectures={sorted(architectures)}" ) - - -def _registry_modules_for_model_cfg(model_cfg: dict): - model_type = str(model_cfg.get("model_type", "")) - - module_names = _MODEL_TYPE_REGISTRY_MODULES.get(model_type) - if module_names is None: - # Leave already-registered plugin/custom models available, but avoid - # importing every built-in module just to produce an unsupported-model - # error. Some built-ins have optional multimodal dependencies. - return () - - if model_type == "qwen" and "visual" in model_cfg: - module_names = module_names + ("lightllm.models.qwen_vl.model",) - elif model_type == "qwen2" and _has_architecture(model_cfg, "RewardModel"): - module_names = module_names + ("lightllm.models.qwen2_reward.model",) - elif model_type == "internlm2" and _has_architecture(model_cfg, "RewardModel"): - module_names = module_names + ("lightllm.models.internlm2_reward.model",) - elif model_type == "llava" and _llava_text_model_type(model_cfg) in {"qwen2", "qwen2_vl", "llama"}: - module_names = module_names + ("lightllm.models.tarsier2.model",) - - return module_names - - -def _ensure_model_registry_bootstrapped(model_cfg: dict) -> None: - module_names = _registry_modules_for_model_cfg(model_cfg) - - for module_name in module_names: - if module_name in _bootstrapped_registry_modules: - continue - import_module(module_name) - _bootstrapped_registry_modules.add(module_name) - return - - -def get_model(model_cfg: dict, model_kvargs: dict): - _ensure_model_registry_bootstrapped(model_cfg) - return _registry_get_model(model_cfg, model_kvargs) - - -def get_model_class(model_cfg: dict): - _ensure_model_registry_bootstrapped(model_cfg) - return _registry_get_model_class(model_cfg) - - -def __getattr__(name): - if name in _MODEL_EXPORTS: - return _load_model_attr(name) - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") - - -__all__ = ["get_model", "get_model_class"] + list(_MODEL_EXPORTS) diff --git a/lightllm/models/deepseek_mtp/model.py b/lightllm/models/deepseek_mtp/model.py index c6ca8dad53..e2b2a56137 100644 --- a/lightllm/models/deepseek_mtp/model.py +++ b/lightllm/models/deepseek_mtp/model.py @@ -6,6 +6,10 @@ class Deepseek3MTPModel(Deepseek2TpPartModel): + + # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). + is_mtp_draft_model = True + pre_and_post_weight_class = Deepseek3MTPPreAndPostLayerWeight pre_layer_infer_class = Deepseek3MTPPreLayerInfer @@ -19,9 +23,6 @@ def _pre_init(self, kvargs: dict): self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") return - def _gen_special_model_input(self, token_num: int): - return self._gen_mtp_draft_special_model_input(token_num) - def _init_custom(self): self._cos_cached = self.main_model._cos_cached self._sin_cached = self.main_model._sin_cached diff --git a/lightllm/models/glm4_moe_lite_mtp/model.py b/lightllm/models/glm4_moe_lite_mtp/model.py index 95ebf23efd..2e4ba5c86b 100644 --- a/lightllm/models/glm4_moe_lite_mtp/model.py +++ b/lightllm/models/glm4_moe_lite_mtp/model.py @@ -9,6 +9,10 @@ class Glm4MoeLiteMTPModel(Glm4MoeLiteTpPartModel): + + # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). + is_mtp_draft_model = True + pre_and_post_weight_class = Glm4MoeLiteMTPPreAndPostLayerWeight pre_layer_infer_class = Deepseek3MTPPreLayerInfer @@ -20,9 +24,6 @@ def _pre_init(self, kvargs: dict): self.main_model: TpPartBaseModel = kvargs.pop("main_model") self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") - def _gen_special_model_input(self, token_num: int): - return self._gen_mtp_draft_special_model_input(token_num) - def _init_custom(self): self._cos_cached = self.main_model._cos_cached self._sin_cached = self.main_model._sin_cached diff --git a/lightllm/models/mistral_mtp/model.py b/lightllm/models/mistral_mtp/model.py index 7c1ac50ba7..f17bc0a383 100644 --- a/lightllm/models/mistral_mtp/model.py +++ b/lightllm/models/mistral_mtp/model.py @@ -9,6 +9,10 @@ class MistralMTPModel(MistralTpPartModel): + + # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). + is_mtp_draft_model = True + pre_and_post_weight_class = MistralMTPPreAndPostLayerWeight pre_layer_infer_class = MistralMTPPreLayerInfer @@ -27,9 +31,6 @@ def _pre_init(self, kvargs: dict): self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") return - def _gen_special_model_input(self, token_num: int): - return self._gen_mtp_draft_special_model_input(token_num) - def _init_some_value(self): super()._init_some_value() self.layers_num = 1 diff --git a/lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py index b22eaeb89d..3d2c613c8e 100644 --- a/lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py +++ b/lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py @@ -9,13 +9,6 @@ class Qwen35DFlashPreAndPostLayerWeight(PreAndPostLayerWeight): - """Qwen3.5 DFlash weights outside the draft decoder stack. - - Qwen3.5 DFlash checkpoints only store the DFlash projection and draft - output norm. Token embedding and LM head weights are shared from the target - model by Qwen3_5DFlashModel. - """ - def __init__(self, data_type, network_config, quant_cfg: Quantcfg): super().__init__(data_type, network_config) self.quant_cfg = quant_cfg @@ -44,4 +37,3 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): weight_name="norm.weight", data_type=self.data_type_, ) - return diff --git a/lightllm/models/qwen3_5_dflash/model.py b/lightllm/models/qwen3_5_dflash/model.py index e27e23dae5..01ca9832f9 100644 --- a/lightllm/models/qwen3_5_dflash/model.py +++ b/lightllm/models/qwen3_5_dflash/model.py @@ -1,4 +1,3 @@ -from lightllm.common.speculative.config import normalize_speculative_draft_config from lightllm.distributed.communication_op import dist_group_manager from lightllm.models.llama.model import LlamaTpPartModel from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel @@ -8,39 +7,26 @@ class Qwen3_5DFlashModel(Qwen3DFlashModel): - """Qwen3.5 DFlash draft model. - - The checkpoint stores DFlash parameters under `dflash_config` and omits - token embedding / LM head weights. Those two weights are shared from the - already-loaded target model. Its query/proposal layout is resolved from the - checkpoint schema by the generic speculative runtime. - """ + """Qwen3.5 DFlash draft model.""" pre_and_post_weight_class = Qwen35DFlashPreAndPostLayerWeight share_target_embedding_and_lm_head = True def _init_config(self): super()._init_config() - self._normalize_dflash_config() - return - - def _normalize_dflash_config(self): - normalize_speculative_draft_config(self.config) + dflash_config = self.config.get("dflash_config", {}) + for key in ("target_layer_ids", "mask_token_id"): + if key not in self.config and key in dflash_config: + self.config[key] = dflash_config[key] rope_parameters = self.config.get("rope_parameters") if isinstance(rope_parameters, dict): if "rope_theta" in rope_parameters and "rope_theta" not in self.config: self.config["rope_theta"] = rope_parameters["rope_theta"] - if ( - "partial_rotary_factor" in rope_parameters - and "partial_rotary_factor" not in self.config - ): - self.config["partial_rotary_factor"] = rope_parameters[ - "partial_rotary_factor" - ] + if "partial_rotary_factor" in rope_parameters and "partial_rotary_factor" not in self.config: + self.config["partial_rotary_factor"] = rope_parameters["partial_rotary_factor"] if "rope_scaling" not in self.config: self.config["rope_scaling"] = rope_parameters - return def _init_custom(self): # The Qwen3.5 target uses mrope/partial rotary with a different head @@ -50,28 +36,19 @@ def _init_custom(self): self.dist_group = dist_group_manager.get_default_group() self.block_size = int(self.config["block_size"]) self.mask_token_id = int(self.config["mask_token_id"]) - return def _init_mem_manager(self): target_mem_manager = self.main_model.mem_manager - target_linear_config = getattr(target_mem_manager, "linear_config", None) - assert ( - target_linear_config is not None - ), "Qwen3_5DFlashModel requires a Qwen3Next target" + target_linear_config = target_mem_manager.linear_config draft_head_dim = self.config.get("head_dim") if draft_head_dim is None: - draft_head_dim = ( - self.config["hidden_size"] // self.config["num_attention_heads"] - ) + draft_head_dim = self.config["hidden_size"] // self.config["num_attention_heads"] draft_kv_heads = int(self.config["num_key_value_heads"]) target_kv_heads = int(target_linear_config.full_att_all_num_kv_heads) draft_layer_num = int(self.config["n_layer"]) reserved_draft_layer_num = int(target_linear_config.draft_full_att_kv_layer_num) - assert ( - int(draft_head_dim) == int(target_mem_manager.head_dim) - and draft_kv_heads == target_kv_heads - ), ( + assert int(draft_head_dim) == int(target_mem_manager.head_dim) and draft_kv_heads == target_kv_heads, ( "Qwen3.5 block draft currently requires draft and target full-attention KV shapes to match: " f"draft=({draft_kv_heads}, {draft_head_dim}), " f"target=({target_kv_heads}, {target_mem_manager.head_dim})" @@ -85,10 +62,5 @@ def _init_mem_manager(self): def _init_weights(self, start_layer_index=None): super()._init_weights(start_layer_index=start_layer_index) if self.share_target_embedding_and_lm_head: - self.pre_post_weight.wte_weight_ = ( - self.main_model.pre_post_weight.wte_weight_ - ) - self.pre_post_weight.lm_head_weight_ = ( - self.main_model.pre_post_weight.lm_head_weight_ - ) - return + self.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ + self.pre_post_weight.lm_head_weight_ = self.main_model.pre_post_weight.lm_head_weight_ diff --git a/lightllm/models/qwen3_5_dspark/model.py b/lightllm/models/qwen3_5_dspark/model.py index 102b8a0d76..bf291d1a76 100644 --- a/lightllm/models/qwen3_5_dspark/model.py +++ b/lightllm/models/qwen3_5_dspark/model.py @@ -2,12 +2,13 @@ from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import ( Qwen3DSparkPostLayerInfer, ) +from lightllm.models.qwen3_dspark.model import DSparkModelOutputMixin from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import ( Qwen3DSparkPreAndPostLayerWeight, ) -class Qwen3_5DSparkModel(Qwen3_5DFlashModel): +class Qwen3_5DSparkModel(DSparkModelOutputMixin, Qwen3_5DFlashModel): """DSpark draft model paired with a Qwen3.5 hybrid-attention target. DeepSpec exports these checkpoints as ``Qwen3DSparkModel`` because the diff --git a/lightllm/models/qwen3_dflash/infer_struct.py b/lightllm/models/qwen3_dflash/infer_struct.py index 6ea456fddf..a391089087 100644 --- a/lightllm/models/qwen3_dflash/infer_struct.py +++ b/lightllm/models/qwen3_dflash/infer_struct.py @@ -1,16 +1,8 @@ -import torch - from lightllm.models.llama.infer_struct import LlamaInferStateInfo class Qwen3DFlashInferStateInfo(LlamaInferStateInfo): - """DFlash metadata on top of the normal prefill state.""" - - def __init__(self): - super().__init__() - self.prefill_causal: bool = True - self.decode_causal: bool = True - self.decode_mtp_step: int = None + """DFlash attention metadata.""" def init_some_extra_state(self, model): super().init_some_extra_state(model) @@ -18,23 +10,3 @@ def init_some_extra_state(self, model): self.prefill_causal = False else: self.decode_causal = False - self.decode_mtp_step = model.block_size - 1 - self.is_draft_model = True - return - - @staticmethod - def build_draft_query_position_ids( - *, - selected_seq_len: torch.Tensor, - b_position_delta: torch.Tensor = None, - draft_step: int, - ) -> torch.Tensor: - offsets = torch.arange( - int(draft_step), - dtype=torch.long, - device=selected_seq_len.device, - ) - position_ids = selected_seq_len.to(torch.long).view(-1, 1) + offsets.view(1, -1) - if b_position_delta is not None: - position_ids = position_ids + b_position_delta.to(torch.long).view(-1, 1) - return position_ids diff --git a/lightllm/models/qwen3_dflash/layer_infer/__init__.py b/lightllm/models/qwen3_dflash/layer_infer/__init__.py index 16e5628d81..e69de29bb2 100644 --- a/lightllm/models/qwen3_dflash/layer_infer/__init__.py +++ b/lightllm/models/qwen3_dflash/layer_infer/__init__.py @@ -1,9 +0,0 @@ -from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer -from lightllm.models.qwen3_dflash.layer_infer.pre_layer_infer import Qwen3DFlashPreLayerInfer -from lightllm.models.qwen3_dflash.layer_infer.transformer_layer_infer import Qwen3DFlashTransformerLayerInfer - -__all__ = [ - "Qwen3DFlashPostLayerInfer", - "Qwen3DFlashPreLayerInfer", - "Qwen3DFlashTransformerLayerInfer", -] diff --git a/lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py index 714513e211..b140c80d5f 100644 --- a/lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py +++ b/lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py @@ -4,23 +4,10 @@ class Qwen3DFlashPostLayerInfer(LlamaPostLayerInfer): - def _is_commit_prefill(self, infer_state): - return infer_state.is_prefill and infer_state.mtp_draft_input_hiddens is not None - - def _tpsp_allgather(self, input: torch.Tensor, infer_state): - if self._is_commit_prefill(infer_state): - return input - return super()._tpsp_allgather(input=input, infer_state=infer_state) - def token_forward(self, input_embdings: torch.Tensor, infer_state, layer_weight): - if self._is_commit_prefill(infer_state): - # Commit prefill only materializes draft KV. There is no LM head - # work to do, but BaseModel still expects a logits-shaped tensor. - return torch.empty( - (infer_state.input_ids.shape[0], 0), - dtype=input_embdings.dtype, - device=input_embdings.device, - ) + if infer_state.is_prefill: + # Commit prefill only writes draft KV; BaseModel still requires a tensor output. + return input_embdings.new_empty((0,)) return super().token_forward( input_embdings=input_embdings, infer_state=infer_state, diff --git a/lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py b/lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py index b7e59529e7..bafe533a91 100644 --- a/lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py +++ b/lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py @@ -1,40 +1,13 @@ -import torch - from lightllm.models.llama.layer_infer.pre_layer_infer import LlamaPreLayerInfer from lightllm.models.qwen3_dflash.layer_weights.pre_and_post_layer_weight import Qwen3DFlashPreAndPostLayerWeight class Qwen3DFlashPreLayerInfer(LlamaPreLayerInfer): - """DFlash target-hidden projection plus normal token embedding.""" + """Project target hiddens for DFlash commit prefill.""" def __init__(self, network_config): super().__init__(network_config) self.eps_ = network_config["rms_norm_eps"] - self.hidden_size_ = network_config["hidden_size"] - return - - def project_target_hidden( - self, - *, - target_hidden_states: torch.Tensor, - layer_weight: Qwen3DFlashPreAndPostLayerWeight, - ) -> torch.Tensor: - if target_hidden_states.dim() == 2: - batch_size = target_hidden_states.shape[0] - context_len = 1 - flat_hidden = target_hidden_states - else: - assert target_hidden_states.dim() == 3 - batch_size, context_len, _ = target_hidden_states.shape - flat_hidden = target_hidden_states.reshape(batch_size * context_len, -1) - - projected = layer_weight.fc_weight_.mm(flat_hidden, use_custom_tensor_mananger=False) - projected = layer_weight.hidden_norm_weight_( - input=projected, - eps=self.eps_, - alloc_func=self.alloc_tensor, - ) - return projected.view(batch_size, context_len, self.hidden_size_) def context_forward( self, @@ -42,26 +15,12 @@ def context_forward( infer_state, layer_weight: Qwen3DFlashPreAndPostLayerWeight, ): - if infer_state.mtp_draft_input_hiddens is None: - return super().context_forward(input_ids, infer_state, layer_weight) - - return self.project_target_hidden( - target_hidden_states=infer_state.mtp_draft_input_hiddens, - layer_weight=layer_weight, - ).reshape(-1, self.hidden_size_) - - def token_forward( - self, - input_ids, - infer_state, - layer_weight: Qwen3DFlashPreAndPostLayerWeight, - ): - return super().token_forward(input_ids, infer_state, layer_weight) - - def decode_forward( - self, - input_ids, - infer_state, - layer_weight: Qwen3DFlashPreAndPostLayerWeight, - ): - return self.token_forward(input_ids, infer_state, layer_weight) + target_hidden_states = layer_weight.fc_weight_.mm( + infer_state.mtp_draft_input_hiddens, + use_custom_tensor_mananger=False, + ) + return layer_weight.hidden_norm_weight_( + input=target_hidden_states, + eps=self.eps_, + alloc_func=self.alloc_tensor, + ) diff --git a/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py index 31a6fba7a7..610ebb6641 100644 --- a/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py @@ -1,21 +1,12 @@ import torch -from lightllm.common.basemodel.attention import AttControl +from lightllm.common.basemodel.triton_kernel.norm.qk_norm import qk_rmsnorm_forward from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo from lightllm.models.qwen3_dflash.layer_weights.transformer_layer_weight import Qwen3DFlashTransformerLayerWeight -def get_draft_layer_type(layer_num, network_config): - layer_types = network_config.get("layer_types", []) - draft_layer_start = int(network_config.get("_draft_layer_start", 0)) - local_layer_num = int(layer_num) - draft_layer_start - if 0 <= local_layer_num < len(layer_types): - return layer_types[local_layer_num] - return "full_attention" - - class Qwen3DFlashTransformerLayerInfer(LlamaTransformerLayerInfer): """DFlash layer inference. @@ -27,13 +18,6 @@ class Qwen3DFlashTransformerLayerInfer(LlamaTransformerLayerInfer): def __init__(self, layer_num, network_config): super().__init__(layer_num, network_config) self.head_dim_ = network_config["head_dim"] - self.block_size_ = int(network_config["block_size"]) - layer_type = get_draft_layer_type(layer_num, network_config) - sliding_window = int(network_config.get("sliding_window", 0) or 0) - self.use_sliding_window_ = bool(network_config.get("use_sliding_window", False)) - self.use_sliding_window_ = self.use_sliding_window_ and layer_type == "sliding_attention" and sliding_window > 0 - self.sliding_window_ = sliding_window - return def context_forward( self, @@ -42,108 +26,41 @@ def context_forward( layer_weight: Qwen3DFlashTransformerLayerWeight, ) -> torch.Tensor: token_num, _ = input_embdings.shape - kv = layer_weight.kv_proj.mm(input_embdings, use_custom_tensor_mananger=False) - kv = kv.view(token_num, self.tp_k_head_num_ + self.tp_v_head_num_, self.head_dim_) - k = kv[:, : self.tp_k_head_num_, :] - v = kv[:, self.tp_k_head_num_ :, :] - k = layer_weight.k_norm_weight_( - input=k.reshape(-1, self.head_dim_), - eps=self.eps_, - alloc_func=torch.empty, - ).view(token_num, self.tp_k_head_num_, self.head_dim_) + cache_kv = layer_weight.kv_proj.mm(input_embdings, use_custom_tensor_mananger=False) + qk_rmsnorm_forward( + cache_kv[:, : self.tp_k_head_num_ * self.head_dim_], + layer_weight.qk_norm_weight_.k_weight, + self.eps_, + ) + cache_kv = cache_kv.view(token_num, self.tp_k_head_num_ + self.tp_v_head_num_, self.head_dim_) rotary_emb_fwd( - k, + cache_kv[:, : self.tp_k_head_num_, :], None, infer_state.position_cos, infer_state.position_sin, ) - cache_kv = torch.cat([k, v], dim=1) - self._post_cache_kv(cache_kv.contiguous(), infer_state, layer_weight) + self._post_cache_kv(cache_kv, infer_state, layer_weight) return input_embdings - def token_forward( - self, - input_embdings: torch.Tensor, - infer_state: Qwen3DFlashInferStateInfo, - layer_weight: Qwen3DFlashTransformerLayerWeight, - ) -> torch.Tensor: - hidden_states = input_embdings.view(-1, self.block_size_, self.embed_dim_) - residual = hidden_states - q, cache_kv = self._get_qkv(hidden_states, infer_state, layer_weight) - batch_size, block_size, hidden_size = hidden_states.shape - self._post_cache_kv(cache_kv.contiguous(), infer_state, layer_weight) - o = self._token_attention_kernel(q, infer_state, layer_weight) - o = self._get_o(o, infer_state=infer_state, layer_weight=layer_weight) - hidden_states = residual + o.view(-1, block_size, hidden_size) - - residual = hidden_states - ffn_input = self._ffn_norm( - hidden_states.reshape(-1, hidden_size), - infer_state=infer_state, - layer_weight=layer_weight, - ) - ffn_out = self._ffn(ffn_input, infer_state=infer_state, layer_weight=layer_weight) - hidden_states = residual + ffn_out.view(-1, block_size, hidden_size) - return hidden_states.view(-1, self.embed_dim_) - def _get_qkv(self, input, infer_state: Qwen3DFlashInferStateInfo, layer_weight: Qwen3DFlashTransformerLayerWeight): - hidden_states = input.view(-1, self.block_size_, self.embed_dim_) - batch_size, block_size, hidden_size = hidden_states.shape - normed = self._att_norm( - hidden_states.reshape(-1, hidden_size), - infer_state=infer_state, - layer_weight=layer_weight, - ).view(batch_size, block_size, hidden_size) - - q = layer_weight.q_proj.mm(normed.reshape(-1, hidden_size), use_custom_tensor_mananger=False) - kv = layer_weight.kv_proj.mm(normed.reshape(-1, hidden_size), use_custom_tensor_mananger=False) - - q = q.view(batch_size, block_size, self.tp_q_head_num_, self.head_dim_) - kv = kv.view(batch_size, block_size, self.tp_k_head_num_ + self.tp_v_head_num_, self.head_dim_) - k = kv[:, :, : self.tp_k_head_num_, :] - v = kv[:, :, self.tp_k_head_num_ :, :] + q = layer_weight.q_proj.mm(input, use_custom_tensor_mananger=False) + cache_kv = layer_weight.kv_proj.mm(input, use_custom_tensor_mananger=False) - q = layer_weight.q_norm_weight_( - input=q.reshape(-1, self.head_dim_), + layer_weight.qk_norm_weight_( + q, + cache_kv[:, : self.tp_k_head_num_ * self.head_dim_], eps=self.eps_, - alloc_func=torch.empty, - ).view(batch_size, block_size, self.tp_q_head_num_, self.head_dim_) - k = layer_weight.k_norm_weight_( - input=k.reshape(-1, self.head_dim_), - eps=self.eps_, - alloc_func=torch.empty, - ).view(batch_size, block_size, self.tp_k_head_num_, self.head_dim_) - - rotary_emb_fwd( - q.reshape(-1, self.tp_q_head_num_, self.head_dim_), - k.reshape(-1, self.tp_k_head_num_, self.head_dim_), - infer_state.position_cos, - infer_state.position_sin, ) - cache_kv = torch.cat([k, v], dim=2).reshape( - batch_size * block_size, + cache_kv = cache_kv.view( + -1, self.tp_k_head_num_ + self.tp_v_head_num_, self.head_dim_, ) - return q, cache_kv - def _token_attention_kernel( - self, - q: torch.Tensor, - infer_state: Qwen3DFlashInferStateInfo, - layer_weight: Qwen3DFlashTransformerLayerWeight, - ) -> torch.Tensor: - _k, _v = infer_state.mem_manager.get_att_input_params(layer_index=self.layer_num_) - _q = q.view(-1, self.tp_q_head_num_, self.head_dim_) - if self.use_sliding_window_: - att_control = AttControl(use_sliding_window=True, sliding_window=(self.sliding_window_ - 1, 0)) - else: - att_control = AttControl() - o_tensor = infer_state.decode_att_state.decode_att( - q=_q, - k=_k, - v=_v, - att_control=att_control, - alloc_func=self.alloc_tensor, + rotary_emb_fwd( + q.view(-1, self.tp_q_head_num_, self.head_dim_), + cache_kv[:, : self.tp_k_head_num_, :], + infer_state.position_cos, + infer_state.position_sin, ) - return o_tensor.view(q.shape) + return q, cache_kv diff --git a/lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py index 69807ceba4..ce2ca006e1 100644 --- a/lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py +++ b/lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py @@ -60,4 +60,3 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): weight_name="lm_head.weight", data_type=self.data_type_, ) - return diff --git a/lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py b/lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py index d7ae61044f..0eaf6b8684 100644 --- a/lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py @@ -1,4 +1,9 @@ -from lightllm.common.basemodel.layer_weights.meta_weights import COLMMWeight, KVROWNMMWeight, RMSNormWeight, ROWMMWeight +from lightllm.common.basemodel.layer_weights.meta_weights import ( + COLMMWeight, + KVROWNMMWeight, + QKRMSNORMWeight, + ROWMMWeight, +) from lightllm.models.llama.layer_weights.transformer_layer_weight import LlamaTransformerLayerWeight @@ -31,7 +36,6 @@ def _init_weight_names(self): self._ffn_norm_weight_name = f"{weight_prefix}.post_attention_layernorm.weight" self._q_norm_name = f"{weight_prefix}.self_attn.q_norm.weight" self._k_norm_name = f"{weight_prefix}.self_attn.k_norm.weight" - return def _init_qkv(self): in_dim = self.n_embed @@ -53,7 +57,6 @@ def _init_qkv(self): bias_names=[self._k_bias_name, self._v_bias_name], quant_method=self.get_quant_method("kv_proj"), ) - return def _init_o(self): in_dim = self.o_head_num_ * self.head_dim @@ -66,7 +69,6 @@ def _init_o(self): bias_names=self._o_bias_name, quant_method=self.get_quant_method("o_proj"), ) - return def _init_ffn(self): self.gate_up_proj = ROWMMWeight( @@ -85,18 +87,12 @@ def _init_ffn(self): bias_names=self._down_bias_name, quant_method=self.get_quant_method("down_proj"), ) - return def _init_norm(self): super()._init_norm() - self.q_norm_weight_ = RMSNormWeight( + self.qk_norm_weight_ = QKRMSNORMWeight( dim=self.head_dim, - weight_name=self._q_norm_name, + q_weight_name=self._q_norm_name, + k_weight_name=self._k_norm_name, data_type=self.data_type_, ) - self.k_norm_weight_ = RMSNormWeight( - dim=self.head_dim, - weight_name=self._k_norm_name, - data_type=self.data_type_, - ) - return diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py index ea01f18251..a34d26a725 100644 --- a/lightllm/models/qwen3_dflash/model.py +++ b/lightllm/models/qwen3_dflash/model.py @@ -1,4 +1,3 @@ - from lightllm.common.basemodel.attention import ( BaseAttBackend, Fa3AttBackend, @@ -7,7 +6,6 @@ get_prefill_att_backend_class, ) from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.cuda_graph import CudaGraph from lightllm.distributed.communication_op import dist_group_manager from lightllm.models.llama.model import LlamaTpPartModel from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo @@ -19,32 +17,9 @@ class Qwen3DFlashModel(LlamaTpPartModel): - """Qwen3 DFlash draft model. - - This is the LightLLM service port of the DeepSpec DFlash/DSpark Qwen3 model. - The service path enters through `forward(ModelInput)` using the same - primitive metadata as normal LightLLM prefill: `mem_indexes`, - `req_to_token_indexs`, sequence lengths, and prefill start locations. - - Target -> draft inputs: - - target hidden rows are committed into DFlash draft K/V. - - this model then materializes the next DFlash block from query/mask - embeddings and scratch KV slots. - - Draft output: - - logits are returned for the flattened [batch, draft_step] block rows. - The proposer maps them back to the standard - [verify_batch, draft_step + 1] speculative proposal shape. - - KV ownership: - - DFlash intentionally reuses `main_model.req_manager` and - `main_model.mem_manager` when the target/draft KV shape is compatible. - The target manager can be over-provisioned with draft layer slots and - extra token capacity so CPU cache, offload, and free logic remain unified. - - The DFlash-specific invariant is the layer/token-slot lifecycle: accepted - target hidden rows become committed draft K/V, while current block K/V is - scratch and must be released or overwritten after proposal/verification. - """ + """Qwen3 DFlash draft model.""" + + is_mtp_draft_model = True pre_and_post_weight_class = Qwen3DFlashPreAndPostLayerWeight transformer_weight_class = Qwen3DFlashTransformerLayerWeight @@ -56,13 +31,14 @@ class Qwen3DFlashModel(LlamaTpPartModel): def __init__(self, kvargs: dict): self._pre_init(kvargs) super().__init__(kvargs) - return def _pre_init(self, kvargs: dict): self.main_model: TpPartBaseModel = kvargs.pop("main_model") self.mtp_previous_draft_models = kvargs.pop("mtp_previous_draft_models") - kvargs["return_all_prompt_logics"] = True - return + + def _verify_params(self): + super()._verify_params() + assert not self.enable_tpsp_mix_mode, "Qwen3 DFlash draft model does not support TP-SP" def _init_custom(self): self._cos_cached = self.main_model._cos_cached @@ -70,19 +46,13 @@ def _init_custom(self): self.dist_group = dist_group_manager.get_default_group() self.block_size = int(self.config["block_size"]) self.mask_token_id = int(self.config["mask_token_id"]) - return def _init_req_manager(self): self.req_manager = self.main_model.req_manager - return def _init_mem_manager(self): - # Intentionally shared with the target model. DFlash uses compatible - # KV shapes, so the main manager should be provisioned with the draft - # layer range and any extra temporary token capacity needed by block - # proposals. This keeps req/cpu-cache/offload/free paths unified. + # Draft KV uses target-owned cache slots. self.mem_manager = self.main_model.mem_manager - return def _init_att_backend(self): self.prefill_att_backend: BaseAttBackend = get_prefill_att_backend_class(index=0)(model=self) @@ -101,7 +71,6 @@ def _init_att_backend(self): "Qwen3DFlashModel requires FA3 decode attention: " "block draft attention is non-causal and Triton/FlashInfer decode paths do not honor decode_causal." ) - return def _init_infer_layer(self, start_layer_index=None): assert start_layer_index is None @@ -109,9 +78,7 @@ def _init_infer_layer(self, start_layer_index=None): self.draft_layer_start += sum( len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models ) - self.config["_draft_layer_start"] = self.draft_layer_start super()._init_infer_layer(start_layer_index=self.draft_layer_start) - return def _init_weights(self, start_layer_index=None): assert start_layer_index is None @@ -129,7 +96,6 @@ def _init_weights(self, start_layer_index=None): ) for i in range(self.config["n_layer"]) ] - return def _gen_special_model_input(self, token_num: int): return {"mtp_draft_input_hiddens": None} @@ -140,21 +106,5 @@ def _autotune_warmup(self): def _init_padded_req(self): return - def _init_cudagraph(self): - if self.disable_cudagraph or self.args.enable_decode_microbatch_overlap: - self.graph = None - return - - self.graph = CudaGraph( - max_batch_size=self.graph_max_batch_size, - max_len_in_batch=self.graph_max_len_in_batch, - tp_world_size=self.tp_world_size_, - ) - return - def _init_prefill_cuda_graph(self): self.prefill_graph = None - return - - def _check_max_len_infer(self): - return diff --git a/lightllm/models/qwen3_dspark/layer_infer/__init__.py b/lightllm/models/qwen3_dspark/layer_infer/__init__.py index 4e412110b5..e69de29bb2 100644 --- a/lightllm/models/qwen3_dspark/layer_infer/__init__.py +++ b/lightllm/models/qwen3_dspark/layer_infer/__init__.py @@ -1,4 +0,0 @@ -from .post_layer_infer import Qwen3DSparkPostLayerInfer - - -__all__ = ["Qwen3DSparkPostLayerInfer"] diff --git a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py index 4a9490359d..5fca8a44b3 100644 --- a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py +++ b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py @@ -25,17 +25,17 @@ def __init__(self, network_config): self.markov_head_type_ = str(network_config.get("markov_head_type", "")).lower() self.enable_confidence_head_ = bool(network_config.get("enable_confidence_head", False)) self.confidence_head_with_markov_ = bool(network_config.get("confidence_head_with_markov", False)) - self.mtp_draft_confidence_logits = None - self.mtp_draft_token_ids = None + self.confidence_logits = None + self.draft_token_ids = None - def pop_mtp_draft_confidence_logits(self): - logits = self.mtp_draft_confidence_logits - self.mtp_draft_confidence_logits = None + def pop_confidence_logits(self): + logits = self.confidence_logits + self.confidence_logits = None return logits - def pop_mtp_draft_token_ids(self): - token_ids = self.mtp_draft_token_ids - self.mtp_draft_token_ids = None + def pop_draft_token_ids(self): + token_ids = self.draft_token_ids + self.draft_token_ids = None return token_ids def has_markov_head(self) -> bool: @@ -69,7 +69,6 @@ def _markov_project_bias( def _markov_step_bias( self, - *, prev_token_ids: torch.Tensor, hidden_states: torch.Tensor, state: torch.Tensor, @@ -100,7 +99,6 @@ def _markov_step_bias( def apply_markov_logits( self, base_logits: torch.Tensor, - *, block_hidden: torch.Tensor, anchor_token_ids: torch.Tensor, layer_weight: Qwen3DSparkPreAndPostLayerWeight, @@ -131,7 +129,6 @@ def apply_markov_logits( def predict_confidence_logits( self, block_hidden: torch.Tensor, - *, anchor_token_ids: torch.Tensor, sampled_tokens: torch.Tensor, layer_weight: Qwen3DSparkPreAndPostLayerWeight, @@ -210,7 +207,6 @@ def _token_forward_with_local_logits_and_hidden( def _sample_tp_sharded_vanilla_markov( self, local_logits: torch.Tensor, - *, infer_state: LlamaInferStateInfo, anchor_token_ids: torch.Tensor, layer_weight: Qwen3DSparkPreAndPostLayerWeight, @@ -239,7 +235,7 @@ def _sample_tp_sharded_vanilla_markov( for step_idx in range(self.block_size_): prev_embeddings = self._markov_prev_embeddings(prev_token_ids, layer_weight) local_markov_bias = F.linear(prev_embeddings.to(dtype=local_markov_w2.dtype), local_markov_w2) - local_base_logits = local_logits[:, step_idx::self.block_size_].permute(1, 0).float() + local_base_logits = local_logits[:, step_idx :: self.block_size_].permute(1, 0).float() local_scores = local_base_logits + local_markov_bias local_max_values, local_max_indexes = torch.max(local_scores, dim=-1) local_token_ids = local_max_indexes + local_start @@ -276,27 +272,17 @@ def token_forward( infer_state: LlamaInferStateInfo, layer_weight: Qwen3DSparkPreAndPostLayerWeight, ): - self.mtp_draft_confidence_logits = None - self.mtp_draft_token_ids = None - if self._is_commit_prefill(infer_state): - return super().token_forward( - input_embdings=input_embdings, - infer_state=infer_state, - layer_weight=layer_weight, - ) - + self.confidence_logits = None + self.draft_token_ids = None if infer_state.is_prefill: - logits, _ = self._token_forward_with_hidden( + return super().token_forward( input_embdings=input_embdings, infer_state=infer_state, layer_weight=layer_weight, ) - return logits use_tp_sharded_markov = ( - self.tp_world_size_ > 1 - and self.has_markov_head() - and self.markov_head_type_ == "vanilla" + self.tp_world_size_ > 1 and self.has_markov_head() and self.markov_head_type_ == "vanilla" ) if use_tp_sharded_markov: local_logits, head_hidden = self._token_forward_with_local_logits_and_hidden( @@ -315,14 +301,14 @@ def token_forward( anchor_token_ids=anchor_token_ids, layer_weight=layer_weight, ) - self.mtp_draft_token_ids = sampled_tokens.reshape(-1) - self.mtp_draft_confidence_logits = self.predict_confidence_logits( + self.draft_token_ids = sampled_tokens.reshape(-1) + self.confidence_logits = self.predict_confidence_logits( block_hidden, anchor_token_ids=anchor_token_ids, sampled_tokens=sampled_tokens, layer_weight=layer_weight, ) - # The proposer consumes mtp_draft_token_ids directly. Keep the + # The proposer consumes draft_token_ids directly. Keep the # leading row dimension for generic graph padding/unpadding while # avoiding an otherwise unused [rows, vocab] tensor. A single # placeholder column is required because CUDA graph's no-ref @@ -349,7 +335,7 @@ def token_forward( anchor_token_ids=anchor_token_ids, layer_weight=layer_weight, ) - self.mtp_draft_confidence_logits = self.predict_confidence_logits( + self.confidence_logits = self.predict_confidence_logits( block_hidden, anchor_token_ids=anchor_token_ids, sampled_tokens=sampled_tokens, diff --git a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py index 3555ad1405..bf63a524ec 100644 --- a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py +++ b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py @@ -64,4 +64,3 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): weight_shape=(1, confidence_input_dim), bias_shape=(1,), ) - return diff --git a/lightllm/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py index 99029c5a1a..48a870a66d 100644 --- a/lightllm/models/qwen3_dspark/model.py +++ b/lightllm/models/qwen3_dspark/model.py @@ -1,9 +1,46 @@ from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer +from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import Qwen3DSparkPreAndPostLayerWeight +from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor -class Qwen3DSparkModel(Qwen3DFlashModel): +class DSparkModelOutputMixin: + def _token_forward(self, infer_state): + model_output = super()._token_forward(infer_state) + confidence_logits = self.post_infer.pop_confidence_logits() + draft_token_ids = self.post_infer.pop_draft_token_ids() + if infer_state.is_cuda_graph: + if confidence_logits is not None: + confidence_logits = tensor_to_no_ref_tensor(confidence_logits) + if draft_token_ids is not None: + draft_token_ids = tensor_to_no_ref_tensor(draft_token_ids) + + return DSparkModelOutput( + logits=model_output.logits, + spec_hidden=model_output.spec_hidden, + confidence_logits=confidence_logits, + draft_token_ids=draft_token_ids, + ) + + def _create_unpad_decode_model_output(self, model_output: DSparkModelOutput, origin_batch_size: int): + padded_batch_size = model_output.logits.shape[0] + model_output = super()._create_unpad_decode_model_output(model_output, origin_batch_size) + if padded_batch_size == origin_batch_size: + return model_output + + if model_output.draft_token_ids is not None: + model_output.draft_token_ids = model_output.draft_token_ids[:origin_batch_size] + if model_output.confidence_logits is not None: + confidence_rows = model_output.confidence_logits.shape[0] + assert padded_batch_size % confidence_rows == 0 + rows_per_confidence = padded_batch_size // confidence_rows + assert origin_batch_size % rows_per_confidence == 0 + model_output.confidence_logits = model_output.confidence_logits[: origin_batch_size // rows_per_confidence] + return model_output + + +class Qwen3DSparkModel(DSparkModelOutputMixin, Qwen3DFlashModel): """Qwen3 DSpark draft model. DSpark keeps the DFlash block backbone. Its extra Markov/confidence heads @@ -13,3 +50,7 @@ class Qwen3DSparkModel(Qwen3DFlashModel): pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight post_layer_infer_class = Qwen3DSparkPostLayerInfer + + def _verify_params(self): + super()._verify_params() + assert self.config.get("enable_confidence_head", False), "DSpark requires enable_confidence_head=true" diff --git a/lightllm/models/qwen3_dspark/model_output.py b/lightllm/models/qwen3_dspark/model_output.py new file mode 100644 index 0000000000..bff766ec21 --- /dev/null +++ b/lightllm/models/qwen3_dspark/model_output.py @@ -0,0 +1,20 @@ +from dataclasses import dataclass +from typing import Optional + +import torch + +from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor + + +@dataclass +class DSparkModelOutput(ModelOutput): + confidence_logits: Optional[torch.Tensor] = None + draft_token_ids: Optional[torch.Tensor] = None + + def to_no_ref_tensor(self): + super().to_no_ref_tensor() + if self.confidence_logits is not None: + self.confidence_logits = tensor_to_no_ref_tensor(self.confidence_logits) + if self.draft_token_ids is not None: + self.draft_token_ids = tensor_to_no_ref_tensor(self.draft_token_ids) diff --git a/lightllm/models/qwen3_eagle/layer_infer/__init__.py b/lightllm/models/qwen3_eagle/layer_infer/__init__.py index 9efa61f978..e69de29bb2 100644 --- a/lightllm/models/qwen3_eagle/layer_infer/__init__.py +++ b/lightllm/models/qwen3_eagle/layer_infer/__init__.py @@ -1,7 +0,0 @@ -from lightllm.models.qwen3_eagle.layer_infer.pre_layer_infer import Qwen3EaglePreLayerInfer -from lightllm.models.qwen3_eagle.layer_infer.transformer_layer_infer import Qwen3EagleTransformerLayerInfer - -__all__ = [ - "Qwen3EaglePreLayerInfer", - "Qwen3EagleTransformerLayerInfer", -] diff --git a/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py index 9fed1ba284..b52a61f945 100644 --- a/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py +++ b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py @@ -9,32 +9,18 @@ class Qwen3EaglePreLayerInfer(LlamaPreLayerInfer): def __init__(self, network_config): super().__init__(network_config) self.hidden_size_ = network_config["hidden_size"] - return - def prepare_mtp_draft_hiddens( + def prepare_spec_draft_hiddens( self, infer_state: InferStateInfo, layer_weight: Qwen3EaglePreAndPostLayerWeight, ) -> None: - # Keep the ModelInput hidden raw for CUDA graph replay; Eagle layers consume this working buffer. - infer_state.eagle_draft_hidden_states = self.project_mtp_draft_hiddens( - infer_state.mtp_draft_input_hiddens, - layer_weight, - ) - return - - def project_mtp_draft_hiddens( - self, - target_hiddens, - layer_weight: Qwen3EaglePreAndPostLayerWeight, - use_custom_tensor_mananger: bool = True, - ): - if target_hiddens is None or target_hiddens.shape[-1] == self.hidden_size_: - return target_hiddens - return layer_weight.fc_weight_.mm( - target_hiddens, - use_custom_tensor_mananger=use_custom_tensor_mananger, - ) + target_hiddens = infer_state.mtp_draft_input_hiddens + # Target verification provides concatenated auxiliary-layer hiddens (N * H). + # Recurrent draft steps feed the previous draft output, which is already H. + if target_hiddens.shape[-1] != self.hidden_size_: + target_hiddens = layer_weight.fc_weight_.mm(target_hiddens) + infer_state.eagle_draft_hidden_states = target_hiddens def context_forward( self, @@ -42,7 +28,7 @@ def context_forward( infer_state: InferStateInfo, layer_weight: Qwen3EaglePreAndPostLayerWeight, ): - self.prepare_mtp_draft_hiddens(infer_state, layer_weight) + self.prepare_spec_draft_hiddens(infer_state, layer_weight) return super().context_forward(input_ids, infer_state, layer_weight) def token_forward( @@ -51,5 +37,5 @@ def token_forward( infer_state: InferStateInfo, layer_weight: Qwen3EaglePreAndPostLayerWeight, ): - self.prepare_mtp_draft_hiddens(infer_state, layer_weight) + self.prepare_spec_draft_hiddens(infer_state, layer_weight) return super().token_forward(input_ids, infer_state, layer_weight) diff --git a/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py index 946a6f8b50..f362c6af8a 100644 --- a/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py @@ -7,29 +7,21 @@ class Qwen3EagleTransformerLayerInfer(LlamaTransformerLayerInfer): - def __init__(self, layer_num, network_config): - super().__init__(layer_num, network_config) - self.head_dim_ = network_config["head_dim"] - return - def _get_qkv(self, input, infer_state: InferStateInfo, layer_weight: Qwen3EagleTransformerLayerWeight): input_part = self._att_norm(input, infer_state, layer_weight) - target_part = layer_weight.hidden_norm_weight_._native_forward( + target_part = layer_weight.hidden_norm_weight_( input=infer_state.eagle_draft_hidden_states, eps=self.eps_, alloc_func=self.alloc_tensor, ) - input = torch.cat([input_part, target_part], dim=-1) - input = input.view(-1, self.embed_dim_ * 2) - input = self._tpsp_allgather(input, infer_state) - q = layer_weight.q_proj.mm(input) - cache_kv = layer_weight.kv_proj.mm(input) - if layer_weight.qk_norm_weight_ is not None: - layer_weight.qk_norm_weight_( - q, - cache_kv[:, : self.tp_k_head_num_ * self.head_dim_], - eps=self.eps_, - ) + qkv_input = torch.cat([input_part, target_part], dim=-1) + q = layer_weight.q_proj.mm(qkv_input) + cache_kv = layer_weight.kv_proj.mm(qkv_input) + layer_weight.qk_norm_weight_( + q, + cache_kv[:, : self.tp_k_head_num_ * self.head_dim_], + eps=self.eps_, + ) cache_kv = cache_kv.view(-1, (self.tp_k_head_num_ + self.tp_v_head_num_), self.head_dim_) rotary_emb_fwd( q.view(-1, self.tp_q_head_num_, self.head_dim_), @@ -37,11 +29,6 @@ def _get_qkv(self, input, infer_state: InferStateInfo, layer_weight: Qwen3EagleT infer_state.position_cos, infer_state.position_sin, ) - - if infer_state.need_dp_prefill_balance: - q = infer_state._all_to_all_unbalance_get(data=q) - cache_kv = infer_state._all_to_all_unbalance_get(data=cache_kv) - return q, cache_kv def context_forward(self, input_embdings, infer_state: InferStateInfo, layer_weight): diff --git a/lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py index 979d6c04d2..12ef341d51 100644 --- a/lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py +++ b/lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py @@ -34,6 +34,11 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): data_type=self.data_type_, ) + # Compressed-vocabulary EAGLE3 stores two mappings with the checkpoint: + # d2t[draft_id] is an offset, so target_id = draft_id + d2t[draft_id]; + # t2d[target_id] is a bool mask indicating whether target_id exists in the draft vocabulary. + # Inference applies d2t after draft argmax; t2d is loaded for checkpoint compatibility but is not read. + # Without these tensors, draft and target token ids are treated as identical. self.d2t_weight_ = None self.t2d_weight_ = None if draft_vocab_size is not None and target_vocab_size is not None: @@ -60,5 +65,3 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): weight_name="embed_tokens.weight", data_type=self.data_type_, ) - - return diff --git a/lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py b/lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py index 26d2ea3345..53b6d5cc0d 100644 --- a/lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py @@ -1,29 +1,12 @@ -from lightllm.common.basemodel.layer_weights.meta_weights.mm_weight.colmm_weight import COLMMWeight from lightllm.common.basemodel.layer_weights.meta_weights.mm_weight.rowmm_weight import KVROWNMMWeight, ROWMMWeight from lightllm.common.basemodel.layer_weights.meta_weights.norm_weight import QKRMSNORMWeight, RMSNormWeight from lightllm.models.llama.layer_weights.transformer_layer_weight import LlamaTransformerLayerWeight -""" -midlayer.hidden_norm.weight [2048] -midlayer.input_layernorm.weight [2048] -midlayer.mlp.down_proj.weight [2048, 6144] -midlayer.mlp.gate_proj.weight [6144, 2048] -midlayer.mlp.up_proj.weight [6144, 2048] -midlayer.post_attention_layernorm.weight [2048] -midlayer.self_attn.k_proj.weight [512, 4096] -midlayer.self_attn.o_proj.weight [2048, 4096] -midlayer.self_attn.q_proj.weight [4096, 4096] -midlayer.self_attn.v_proj.weight [512, 4096] -""" - class Qwen3EagleTransformerLayerWeight(LlamaTransformerLayerWeight): def _init_weight_names(self): super()._init_weight_names() - if self.network_config_["architectures"][0] in ["Eagle3Speculator", "Qwen3Eagle3Model"]: - weight_prefix = f"layers.{self.layer_num_}" - else: - weight_prefix = "midlayer" + weight_prefix = f"layers.{self.layer_num_}" self._q_weight_name = f"{weight_prefix}.self_attn.q_proj.weight" self._k_weight_name = f"{weight_prefix}.self_attn.k_proj.weight" self._v_weight_name = f"{weight_prefix}.self_attn.v_proj.weight" @@ -62,36 +45,6 @@ def _init_qkv(self): quant_method=self.get_quant_method("kv_proj"), ) - def _init_o(self): - in_dim = self.o_head_num_ * self.head_dim - out_dim = self.n_embed - self.o_proj = COLMMWeight( - in_dim=in_dim, - out_dims=[out_dim], - weight_names=self._o_weight_name, - data_type=self.data_type_, - bias_names=self._o_bias_name, - quant_method=self.get_quant_method("o_proj"), - ) - - def _init_ffn(self): - self.gate_up_proj = ROWMMWeight( - in_dim=self.n_embed, - out_dims=[self.n_inter, self.n_inter], - weight_names=[self._gate_weight_name, self._up_weight_name], - data_type=self.data_type_, - bias_names=[self._gate_bias_name, self._up_bias_name], - quant_method=self.get_quant_method("gate_up_proj"), - ) - self.down_proj = COLMMWeight( - in_dim=self.n_inter, - out_dims=[self.n_embed], - weight_names=self._down_weight_name, - data_type=self.data_type_, - bias_names=self._down_bias_name, - quant_method=self.get_quant_method("down_proj"), - ) - def _init_norm(self): super()._init_norm() hidden_size = self.network_config_["hidden_size"] @@ -100,12 +53,9 @@ def _init_norm(self): weight_name=self._hidden_norm_weight_name, data_type=self.data_type_, ) - self.qk_norm_weight_ = None - architecture = (self.network_config_.get("architectures") or [""])[0] - if architecture in {"Eagle3Speculator", "Qwen3Eagle3Model"}: - self.qk_norm_weight_ = QKRMSNORMWeight( - dim=self.head_dim, - q_weight_name=self._q_norm_name, - k_weight_name=self._k_norm_name, - data_type=self.data_type_, - ) + self.qk_norm_weight_ = QKRMSNORMWeight( + dim=self.head_dim, + q_weight_name=self._q_norm_name, + k_weight_name=self._k_norm_name, + data_type=self.data_type_, + ) diff --git a/lightllm/models/qwen3_eagle/model.py b/lightllm/models/qwen3_eagle/model.py index e6d012da37..93c6deddfd 100644 --- a/lightllm/models/qwen3_eagle/model.py +++ b/lightllm/models/qwen3_eagle/model.py @@ -10,6 +10,8 @@ class Qwen3EagleModel(LlamaTpPartModel): + is_mtp_draft_model = True + pre_and_post_weight_class = Qwen3EaglePreAndPostLayerWeight pre_layer_infer_class = Qwen3EaglePreLayerInfer @@ -19,42 +21,39 @@ class Qwen3EagleModel(LlamaTpPartModel): def __init__(self, kvargs: dict): self._pre_init(kvargs) super().__init__(kvargs) - return def _pre_init(self, kvargs: dict): self.main_model: TpPartBaseModel = kvargs.pop("main_model") self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") - return - def _gen_special_model_input(self, token_num: int): - return self._gen_mtp_draft_special_model_input(token_num) + def _init_config(self): + super()._init_config() + draft_vocab_size = self.config.get("draft_vocab_size") + if draft_vocab_size is not None: + # Inherited model setup reads vocab_size, while EAGLE3's vocabulary + # maps still need the original target vocabulary size. + self.config["target_vocab_size"] = self.config["vocab_size"] + self.config["vocab_size"] = draft_vocab_size + + def _verify_params(self): + super()._verify_params() + assert not self.enable_tpsp_mix_mode, "Qwen3 Eagle draft model does not support TP-SP" def _init_custom(self): self._cos_cached = self.main_model._cos_cached self._sin_cached = self.main_model._sin_cached - return def _init_req_manager(self): self.req_manager = self.main_model.req_manager - return def _init_mem_manager(self): self.mem_manager = self.main_model.mem_manager - return def _init_weights(self, start_layer_index=None): assert start_layer_index is None self.pre_post_weight = self.pre_and_post_weight_class( self.data_type, network_config=self.config, quant_cfg=self.quant_cfg ) - # Compressed-vocabulary EAGLE3 checkpoints store a draft lm_head and - # vocabulary maps, but their input ids still belong to the target - # vocabulary. Reuse the target embedding instead of requiring a draft - # embed_tokens.weight that these checkpoints do not contain. - if self.config.get("draft_vocab_size") is not None: - target_embedding = getattr(self.main_model.pre_post_weight, "wte_weight_", None) - assert target_embedding is not None, "compressed-vocab EAGLE3 requires target token embeddings" - self.pre_post_weight.wte_weight_ = target_embedding self.trans_layers_weight = [ self.transformer_weight_class( i, @@ -64,7 +63,12 @@ def _init_weights(self, start_layer_index=None): ) for i in range(self.config["n_layer"]) ] - return + # Compressed-vocabulary EAGLE3 checkpoints store a draft lm_head and + # vocabulary maps, but their input ids still belong to the target + # vocabulary. Reuse the target embedding instead of requiring a draft + # embed_tokens.weight that these checkpoints do not contain. + if self.config.get("draft_vocab_size") is not None: + self.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ def _init_infer_layer(self, start_layer_index=None): assert start_layer_index is None @@ -73,7 +77,6 @@ def _init_infer_layer(self, start_layer_index=None): len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models ) super()._init_infer_layer(start_layer_index=total_pre_layers_num) - return # d2t stores per-token offsets: target_id = draft_id + d2t[draft_id]. @torch.no_grad() diff --git a/lightllm/models/qwen3_moe_mtp/model.py b/lightllm/models/qwen3_moe_mtp/model.py index f91d9a47e7..d9854250e2 100644 --- a/lightllm/models/qwen3_moe_mtp/model.py +++ b/lightllm/models/qwen3_moe_mtp/model.py @@ -8,6 +8,10 @@ class Qwen3MOEMTPModel(Qwen3MOEModel): + + # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). + is_mtp_draft_model = True + pre_and_post_weight_class = Qwen3MOEMTPPreAndPostLayerWeight pre_layer_infer_class = Deepseek3MTPPreLayerInfer @@ -24,9 +28,6 @@ def _pre_init(self, kvargs: dict): self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") return - def _gen_special_model_input(self, token_num: int): - return self._gen_mtp_draft_special_model_input(token_num) - def _init_custom(self): self._cos_cached = self.main_model._cos_cached self._sin_cached = self.main_model._sin_cached diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 3526980efc..515cacf699 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -613,34 +613,6 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: a new CUDA graph will be generated for every increment of graph_grow_step_size. """, ) - parser.add_argument( - "--mtp_draft_graph_max_batch_size", - type=int, - default=None, - help=""" - Optional logical CUDA graph batch-size limit for MTP draft models. - Defaults to graph_max_batch_size. The value is expanded by the draft - model's decode graph group size in the same way as the main setting. - """, - ) - parser.add_argument( - "--mtp_draft_graph_split_batch_size", - type=int, - default=None, - help=""" - Optional dense-prefix CUDA graph limit for MTP draft models. - Defaults to graph_split_batch_size. - """, - ) - parser.add_argument( - "--mtp_draft_graph_grow_step_size", - type=int, - default=None, - help=""" - Optional CUDA graph batch-size growth step for MTP draft models. - Defaults to graph_grow_step_size. - """, - ) parser.add_argument( "--graph_max_len_in_batch", type=int, @@ -770,7 +742,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: ], default=None, help="""Speculative decoding mode. - *_with_att and *_no_att select attention or non-attention MTP drafts; + *_with_att and *_no_att select attention or non-attention draft models; eagle3 uses a recurrent EAGLE3 draft; dspark and dflash use block draft models.""", ) parser.add_argument( @@ -778,8 +750,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: type=str, nargs="+", default=None, - help="""Path to the draft model for the MTP multi-prediction feature, - used for loading the MTP multi-output token model.""", + help="""Path to the speculative draft model. The legacy option name is + retained for command-line compatibility.""", ) parser.add_argument( "--mtp_step", @@ -791,7 +763,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: parser.add_argument( "--mtp_dynamic_verify", action="store_true", - help="""Whether to enable dynamic verification for MTP multi-prediction results.""", + help="""Compatibility switch for dynamic speculative scheduling. + DSpark enables this scheduling mode automatically.""", ) parser.add_argument( "--kv_quant_calibration_config_path", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index c19096c7f5..16fcea1525 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -3,8 +3,6 @@ import uuid import subprocess import math -from dataclasses import replace -from transformers.configuration_utils import PretrainedConfig from lightllm.utils.start_utils import process_manager from .metrics.manager import start_metric_manager from .embed_cache.manager import start_cache_manager @@ -28,43 +26,10 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args -from lightllm.common.speculative import ( - SpeculativeConfig, - get_block_draft_layout, -) logger = init_logger(__name__) -def normalize_block_mtp_step_from_first_draft_config( - args: StartArgs, spec_config: SpeculativeConfig -) -> SpeculativeConfig: - if not spec_config.uses_block_draft_model: - return spec_config - - assert args.mtp_draft_model_dir is not None and len(args.mtp_draft_model_dir) > 0 - mtp_model_cfg, _ = PretrainedConfig.get_config_dict(args.mtp_draft_model_dir[0]) - layout = get_block_draft_layout( - mtp_model_cfg, - mode=spec_config.mode, - require_confidence_head=spec_config.is_dspark, - ) - configured_step = int(args.mtp_step) - draft_step = layout.resolve_draft_step(configured_step) - if configured_step not in (0, draft_step): - logger.warning( - "Overriding mtp_step=%s with draft_step=%s from block_size=%s for %s mode", - configured_step, - draft_step, - layout.query_block_size, - spec_config.mode, - ) - args.mtp_step = draft_step - spec_config = replace(spec_config, step=draft_step) - spec_config.validate() - return spec_config - - def _set_envs_and_config(args: StartArgs): mp.set_start_method("spawn", force=True) @@ -132,9 +97,9 @@ def _launch_subprocesses(args: StartArgs): args.embed_cache_storage_size = 0.8 args.graph_max_batch_size = 6 logger.info( - "performance_mode is personal, set running_max_req_size to 3," - "batch_max_tokens to 2048, chunked_prefill_size to 1024," - "graph_max_batch_size to 32" + f"performance_mode is personal, set running_max_req_size to 3," + f"batch_max_tokens to 2048, chunked_prefill_size to 1024," + f"graph_max_batch_size to 32" ) if not args.disable_shm_warning: @@ -197,19 +162,30 @@ def _launch_subprocesses(args: StartArgs): ) # mtp params check - spec_config = SpeculativeConfig.from_args(args) - spec_config.validate() - if spec_config.enabled: + spec_mode = args.mtp_mode + if spec_mode is not None: + if spec_mode in ("vanilla_with_att", "vanilla_no_att", "qwen3next_vanilla"): + assert args.mtp_step > 0 + draft_model_count = args.mtp_step + elif spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3", "qwen3next_eagle"): + assert args.mtp_step > 0 + draft_model_count = 1 + else: + assert spec_mode in ("dspark", "dflash"), f"unsupported speculative mode {spec_mode}" + assert args.mtp_step > 0 + draft_model_count = 1 + + if spec_mode == "dspark": + args.mtp_dynamic_verify = True + elif spec_mode == "dflash": + args.mtp_dynamic_verify = False + if args.mtp_draft_model_dir is None: - assert not spec_config.uses_block_draft_model, ( - f"--mtp_draft_model_dir is required for {spec_config.mode} mode" - ) - args.mtp_draft_model_dir = [args.model_dir] * spec_config.draft_model_count - elif isinstance(args.mtp_draft_model_dir, str): - args.mtp_draft_model_dir = [args.mtp_draft_model_dir] - assert len(args.mtp_draft_model_dir) >= spec_config.draft_model_count - spec_config = normalize_block_mtp_step_from_first_draft_config(args, spec_config) + assert spec_mode not in ("dspark", "dflash"), f"--mtp_draft_model_dir is required for {spec_mode} mode" + args.mtp_draft_model_dir = [args.model_dir] * draft_model_count + assert len(args.mtp_draft_model_dir) >= draft_model_count else: + assert args.mtp_step == 0 assert args.mtp_draft_model_dir is None # automatically set visual_dp based on visual_tp and tp. diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index b0a9b59563..39177b1498 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -1,3 +1,4 @@ +import os import math import ctypes import asyncio diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index d5c5a8dd06..af723753a0 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -152,9 +152,6 @@ class StartArgs: graph_max_batch_size: int = field(default=256) graph_split_batch_size: int = field(default=32) graph_grow_step_size: int = field(default=16) - mtp_draft_graph_max_batch_size: Optional[int] = field(default=None) - mtp_draft_graph_split_batch_size: Optional[int] = field(default=None) - mtp_draft_graph_grow_step_size: Optional[int] = field(default=None) graph_max_len_in_batch: int = field(default=0) quant_type: Optional[str] = field(default="none") quant_cfg: Optional[str] = field(default=None) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 94dccaafd0..e87a41529f 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -27,7 +27,7 @@ from lightllm.server.core.objs.out_token_circlequeue import LIGHTLLM_OUT_TOKEN_QUEUE_SIZE from lightllm.server.core.objs.io_objs import GroupReqObjs from lightllm.server.core.objs.shm_req_manager import ShmReqManager -from lightllm.server.core.objs.atomic_array_lock import AtomicShmArrayLock, AsyncLock +from lightllm.server.core.objs.atomic_array_lock import AtomicShmArrayLock, AsyncLock, AtomicLockItem from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt from lightllm.utils.log_utils import init_logger from lightllm.server.metrics.manager import MetricClient @@ -38,7 +38,6 @@ from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.shm_port_args import get_shm_port_args from lightllm.utils.error_utils import ClientDisconnected, PDPrefillNodeStopGenToken -from lightllm.common.speculative import SpeculativeConfig from rpyc.utils.classic import obtain logger = init_logger(__name__) @@ -711,9 +710,6 @@ async def _wait_to_token_package( sub_req_id_to_mtp_accepted_token_num: Dict[int, int] = {} sub_req_id_to_mtp_verify_token_num: Dict[int, int] = {} sub_req_id_to_mtp_verify_step_num: Dict[int, int] = {} - mtp_verify_step_num = 0 - spec_config = SpeculativeConfig.from_args(self.args) - is_static_mtp = spec_config.step > 0 and not spec_config.dynamic_verify first_token_cost_ms = sys.float_info.max prompt_tokens = len(prompt_ids) is_first_token = True @@ -748,14 +744,9 @@ async def _wait_to_token_package( metadata["prompt_cache_len"] = gpu_prompt_cache_len + cpu_prompt_cache_len + disk_prompt_cache_len sub_req_id_to_mtp_accepted_token_num[sub_req_id] = metadata.get("mtp_accepted_token_num", 0) cur_mtp_verify_token_num = metadata.get("mtp_verify_token_num", 0) - prev_mtp_verify_token_num = sub_req_id_to_mtp_verify_token_num.get(sub_req_id, 0) sub_req_id_to_mtp_verify_token_num[sub_req_id] = cur_mtp_verify_token_num cur_mtp_verify_step_num = metadata.get("mtp_verify_step_num", 0) sub_req_id_to_mtp_verify_step_num[sub_req_id] = cur_mtp_verify_step_num - if is_static_mtp and cur_mtp_verify_token_num > prev_mtp_verify_token_num: - mtp_verify_step_num += (cur_mtp_verify_token_num - prev_mtp_verify_token_num) // ( - self.args.mtp_step + 1 - ) if is_first_token: first_token_cost_ms = (time.time() - start_time) * 1000 @@ -773,8 +764,7 @@ async def _wait_to_token_package( unfinished_count -= 1 if unfinished_count == 0: - finish_time = time.time() - total_cost_time_ms = (finish_time - start_time) * 1000 + total_cost_time_ms = (time.time() - start_time) * 1000 mean_per_token_cost_time_ms = (total_cost_time_ms - first_token_cost_ms) / out_token_counter self.per_token_costs.add(mean_per_token_cost_time_ms) x_request_id = request.headers.get("X-Request-Id", "") if request is not None else "" @@ -788,30 +778,18 @@ async def _wait_to_token_package( mtp_accepted_token_num = sum(sub_req_id_to_mtp_accepted_token_num.values()) mtp_verify_token_num = sum(sub_req_id_to_mtp_verify_token_num.values()) - mtp_total_step = sum(sub_req_id_to_mtp_verify_step_num.values()) - if mtp_total_step <= 0: - mtp_total_step = out_token_counter - mtp_accepted_token_num - if mtp_total_step <= 0 and is_static_mtp and mtp_verify_step_num > 0: - mtp_total_step = mtp_verify_step_num - mtp_avg_token_per_step = out_token_counter / max(mtp_total_step, 1) - mtp_avg_verify_tokens_per_step = mtp_verify_token_num / max(mtp_total_step, 1) - mtp_avg_accept_len_per_step_direct = mtp_accepted_token_num / max(mtp_total_step, 1) - decode_start_time = start_time + first_token_cost_ms / 1000.0 - decode_end_time = finish_time - decode_total_time_ms = max((decode_end_time - decode_start_time) * 1000, 0.0) - decode_token_counter = max(out_token_counter - 1, 0) - decode_token_throughput = decode_token_counter / max(decode_total_time_ms / 1000.0, 1e-6) + mtp_total_verify_steps = sum(sub_req_id_to_mtp_verify_step_num.values()) + if mtp_total_verify_steps <= 0: + mtp_total_verify_steps = out_token_counter - mtp_accepted_token_num + mtp_avg_token_per_step = out_token_counter / max(mtp_total_verify_steps, 1) + mtp_avg_verify_tokens_per_step = mtp_verify_token_num / max(mtp_total_verify_steps, 1) + mtp_avg_accepted_tokens_per_step = mtp_accepted_token_num / max(mtp_total_verify_steps, 1) format_start_time = datetime.datetime.fromtimestamp(start_time).strftime("%Y-%m-%d %H:%M:%S") logger.info( f"X-Request-Id:{x_request_id} " f"X-Session-Id:{x_session_id} start_time:{format_start_time} " f"lightllm_req_id:{group_request_id} first_token_cost:{first_token_cost_ms}ms " f"total_cost_time:{total_cost_time_ms}ms,out_token_counter:{out_token_counter} " - f"decode_start_time:{decode_start_time} " - f"decode_end_time:{decode_end_time} " - f"decode_total_time:{decode_total_time_ms}ms " - f"decode_token_counter:{decode_token_counter} " - f"decode_token_throughput:{decode_token_throughput} " f"mean_per_token_cost_time: {mean_per_token_cost_time_ms}ms " f"prompt_token_num:{prompt_tokens} " f"gpu cache hit: {gpu_prompt_cache_ratio > 0} " @@ -824,10 +802,10 @@ async def _wait_to_token_package( f"disk_prompt_cache_len:{disk_prompt_cache_len} " f"disk_prompt_cache_ratio:{disk_prompt_cache_ratio} " f"mtp_accepted_token_num:{mtp_accepted_token_num} " - f"mtp_total_step:{mtp_total_step} " + f"mtp_total_verify_steps:{mtp_total_verify_steps} " f"mtp_total_verify_tokens:{mtp_verify_token_num} " f"mtp_avg_token_per_step:{mtp_avg_token_per_step} " - f"mtp_avg_accept_len_per_step_direct:{mtp_avg_accept_len_per_step_direct} " + f"mtp_avg_accepted_tokens_per_step:{mtp_avg_accepted_tokens_per_step} " f"mtp_avg_verify_tokens_per_step:{mtp_avg_verify_tokens_per_step} " ) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index c3ee798f06..2f1940fde1 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -4,6 +4,7 @@ import uvloop import time import datetime +import ujson as json import pickle import httpx from contextlib import aclosing @@ -423,10 +424,10 @@ async def _wait_to_token_package( x_request_id = request.headers.get("X-Request-Id", "") x_session_id = request.headers.get("X-Session-Id", "") prompt_cache_ratio = prompt_cache_len / prompt_tokens - mtp_total_step = sum(sub_req_id_to_mtp_verify_step_num.values()) - if mtp_total_step <= 0: - mtp_total_step = out_token_counter - sum(sub_req_id_to_mtp_accepted_token_num.values()) - mtp_avg_token_per_step = out_token_counter / max(mtp_total_step, 1) + mtp_total_verify_steps = sum(sub_req_id_to_mtp_verify_step_num.values()) + if mtp_total_verify_steps <= 0: + mtp_total_verify_steps = out_token_counter - sum(sub_req_id_to_mtp_accepted_token_num.values()) + mtp_avg_token_per_step = out_token_counter / max(mtp_total_verify_steps, 1) format_start_time = datetime.datetime.fromtimestamp(start_time).strftime("%Y-%m-%d %H:%M:%S") logger.info( f"X-Request-Id:{x_request_id} " diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 7c350bcc54..56d1c1445e 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -1,20 +1,23 @@ import enum import torch import torch.distributed as dist +import numpy as np import collections import pickle from sortedcontainers import SortedDict -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import List, Dict, Tuple, Optional, Callable, Any, Union from lightllm.common.req_manager import ReqManager, ReqManagerForMamba -from lightllm.server.core.objs import Req, FinishStatus, ShmReqManager +from lightllm.utils.infer_utils import mark_start, mark_end +from lightllm.server.core.objs import Req, SamplingParams, FinishStatus, ShmReqManager from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache, TreeNode from lightllm.server.router.dynamic_prompt.linear_att_radix_cache import ( LinearAttPagedRadixCache, LinearAttPagedTreeNode, ) from lightllm.utils.log_utils import init_logger +from lightllm.server.req_id_generator import convert_sub_id_to_group_id from lightllm.server.multimodal_params import MultimodalParams from lightllm.utils.custom_kernel_utis import custom_cat from lightllm.utils.envs_utils import get_env_start_args @@ -34,7 +37,6 @@ class InferenceContext: infer_req_ids = None vocab_size = None cpu_embed_cache_client: Optional[CpuEmbedCacheClient] = None - dynamic_mtp_planner: Optional[Any] = None overlap_stream: torch.cuda.Stream = None # 一些情况下推理进程进行异步折叠操作的异步流对象。 cpu_kv_cache_stream: torch.cuda.Stream = None # 用 cpu kv cache 操作的 stream @@ -70,47 +72,6 @@ def init_cpu_embed_cache_client(self): self.cpu_embed_cache_client = CpuEmbedCacheClient(create_meta_data=False, init_shm_data=False) return - def init_dynamic_mtp_planner(self, mtp_step: int, mode: str = None): - if mode == "dspark": - planner_mode = "dspark" - elif mode == "eagle3": - planner_mode = "eagle3" - else: - planner_mode = "default" - if ( - self.dynamic_mtp_planner is not None - and self.dynamic_mtp_planner.mtp_step == mtp_step - and getattr(self.dynamic_mtp_planner, "planner_mode", "default") == planner_mode - ): - return - - from lightllm.server.router.model_infer.speculative.planner import ( - DSparkDynamicMTPPlanner, - DynamicMTPPlanner, - Eagle3DynamicMTPPlanner, - ) - - planner_cls = { - "default": DynamicMTPPlanner, - "dspark": DSparkDynamicMTPPlanner, - "eagle3": Eagle3DynamicMTPPlanner, - }[planner_mode] - self.dynamic_mtp_planner = planner_cls(mtp_step=mtp_step) - return - - def record_dynamic_mtp_infer_cost(self, *, batch_size: int, infer_cost_ms: float, is_draft_model: bool): - if self.dynamic_mtp_planner is None: - self.init_dynamic_mtp_planner( - mtp_step=get_env_start_args().mtp_step, - mode=get_env_start_args().mtp_mode, - ) - self.dynamic_mtp_planner.update_infer_cost( - batch_size=batch_size, - infer_cost_ms=infer_cost_ms, - is_draft_model=is_draft_model, - ) - return - def get_overlap_stream(self) -> torch.cuda.Stream: if self.overlap_stream is None: self.overlap_stream = torch.cuda.Stream() @@ -436,7 +397,6 @@ def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: L copy_linear_att_state_to_kv_buffer( b_req_idx=b_req_idx, - req_to_mtp_state_index=self.req_manager.req_to_mtp_state_index, big_page_buffer_ids=big_page_buffer_ids, gpu_conv_state=self.req_manager.req_to_conv_state.buffer, gpu_ssm_state=self.req_manager.req_to_ssm_state.buffer, @@ -606,12 +566,11 @@ def __init__( # 卸载到 cpu cache 中,该标志变量用于标记请求的卸载任务的状态 self.cpu_cache_task_status: "InferReq._CpuCacheTaskStatus" = InferReq._CpuCacheTaskStatus.NOT_STARTED - # mtp_step 用来记录一个请求 draft模型每步需要生成的token数量 + # max_draft_step 用来记录一个请求 draft模型每步需要生成的token数量 # 正常模式下,这个值为0,在 mtp 模式下,这个值为 draft 模型每步需要生成的token数量 - self.mtp_step: int = get_env_start_args().mtp_step - - if self.mtp_step > 0: - self.decode_need_token_num = self._mtp_decode_need_token_num + self.max_draft_step: int = get_env_start_args().mtp_step + if self.max_draft_step > 0: + self.decode_need_token_num = self._spec_decode_need_token_num else: self.decode_need_token_num = self._normal_decode_need_token_num @@ -896,16 +855,14 @@ def set_next_gen_token_id(self, next_token_id: int, logprob: float, output_len: self.shm_req.shm_logprobs.arr[index - 1] = (logprob, rank) return - def update_mtp_accepted_token_num(self, accept_token_num: int): + def update_spec_accepted_token_num(self, accept_token_num: int): # 用于统计 mtp 的接受率 self.shm_req.mtp_accepted_token_num += accept_token_num - def update_mtp_verify_token_num(self, verify_token_num: int): - # 用于统计 mtp 验证时发送给主模型的 token 总数 + def update_spec_verify_token_num(self, verify_token_num: int): self.shm_req.mtp_verify_token_num += verify_token_num - def update_mtp_verify_step_num(self, verify_step_num: int): - # 用于统计 mtp 验证轮数 + def update_spec_verify_step_num(self, verify_step_num: int): self.shm_req.mtp_verify_step_num += verify_step_num def get_last_gen_token(self): @@ -966,8 +923,8 @@ def decode_need_token_num(self) -> int: def _normal_decode_need_token_num(self) -> int: return 1 - def _mtp_decode_need_token_num(self) -> int: - return (1 + self.mtp_step) * 2 + def _spec_decode_need_token_num(self) -> int: + return (1 + self.max_draft_step) * 2 class InferReqUpdatePack: diff --git a/lightllm/server/router/model_infer/mode_backend/__init__.py b/lightllm/server/router/model_infer/mode_backend/__init__.py index f26b329856..8c608e5f92 100644 --- a/lightllm/server/router/model_infer/mode_backend/__init__.py +++ b/lightllm/server/router/model_infer/mode_backend/__init__.py @@ -1,46 +1,15 @@ -from importlib import import_module - - -_BACKEND_EXPORTS = { - "ChunkedPrefillBackend": (".chunked_prefill.impl", "ChunkedPrefillBackend"), - "FirstTokenConstraintBackend": ( - ".chunked_prefill.impl_for_first_token_constraint_mode", - "FirstTokenConstraintBackend", - ), - "OutlinesConstraintBackend": ( - ".chunked_prefill.impl_for_outlines_constraint_mode", - "OutlinesConstraintBackend", - ), - "ReturnPromptLogProbBackend": ( - ".chunked_prefill.impl_for_return_all_prompt_logprobs", - "ReturnPromptLogProbBackend", - ), - "RewardModelBackend": (".chunked_prefill.impl_for_reward_model", "RewardModelBackend"), - "TokenHealingBackend": (".chunked_prefill.impl_for_token_healing", "TokenHealingBackend"), - "XgrammarBackend": (".chunked_prefill.impl_for_xgrammar_mode", "XgrammarBackend"), - "DPChunkedPrefillBackend": (".dp_backend.impl", "DPChunkedPrefillBackend"), - "DiversehBackend": (".diverse_backend.impl", "DiversehBackend"), - "PDChunkedPrefillForPrefillNode": ( - ".pd.prefill_node_impl.prefill_impl", - "PDChunkedPrefillForPrefillNode", - ), - "PDDPChunkedForPrefillNode": ( - ".pd.prefill_node_impl.prefill_impl_for_dp", - "PDDPChunkedForPrefillNode", - ), - "PDDecodeNode": (".pd.decode_node_impl.decode_impl", "PDDecodeNode"), - "PDDPForDecodeNode": (".pd.decode_node_impl.decode_impl_for_dp", "PDDPForDecodeNode"), -} - - -def __getattr__(name): - if name not in _BACKEND_EXPORTS: - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") - - module_name, attr_name = _BACKEND_EXPORTS[name] - value = getattr(import_module(module_name, __name__), attr_name) - globals()[name] = value - return value - - -__all__ = list(_BACKEND_EXPORTS) +from .chunked_prefill.impl import ChunkedPrefillBackend +from .chunked_prefill.impl_for_first_token_constraint_mode import FirstTokenConstraintBackend +from .chunked_prefill.impl_for_outlines_constraint_mode import OutlinesConstraintBackend +from .chunked_prefill.impl_for_reward_model import RewardModelBackend +from .chunked_prefill.impl_for_token_healing import TokenHealingBackend +from .chunked_prefill.impl_for_xgrammar_mode import XgrammarBackend + +from .dp_backend.impl import DPChunkedPrefillBackend +from .diverse_backend.impl import DiversehBackend + +# pd mode backend +from .pd.prefill_node_impl.prefill_impl import PDChunkedPrefillForPrefillNode +from .pd.prefill_node_impl.prefill_impl_for_dp import PDDPChunkedForPrefillNode +from .pd.decode_node_impl.decode_impl import PDDecodeNode +from .pd.decode_node_impl.decode_impl_for_dp import PDDPForDecodeNode 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 926baabff7..6841eff919 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -1,16 +1,16 @@ import os +from collections import Counter + import numpy as np import torch import time import threading import torch.distributed as dist -import collections -from dataclasses import replace -from typing import List, Tuple, Callable, Optional +from typing import List, Tuple, Callable, Optional, Union from transformers.configuration_utils import PretrainedConfig from lightllm.utils.infer_utils import set_random_seed from lightllm.utils.log_utils import init_logger -from lightllm.models import get_model +from lightllm.models import get_draft_model_class, get_model from lightllm.server.router.model_infer.infer_batch import InferReq, InferReqUpdatePack from lightllm.server.router.token_load import TokenLoad from lightllm.common.basemodel.basemodel import TpPartBaseModel @@ -35,20 +35,9 @@ get_env_start_args, enable_radix_tree_timer_merge, get_radix_tree_merge_update_delta, - enable_dynamic_mtp_verify, ) from lightllm.distributed import dist_group_manager -from lightllm.common.speculative import ( - SpeculativeConfig, - get_block_draft_layout, - is_dspark_draft_config, - is_eagle3_draft_config, - is_gemma4_dspark_draft_config, - is_qwen3_dflash_draft_config, - is_qwen3_5_dflash_draft_config, - is_qwen3_dspark_draft_config, -) -from lightllm.server.router.model_infer.speculative import build_spec_runtime +from lightllm.server.router.model_infer.speculative import SpecEngine from lightllm.distributed.communication_op import ( all_gather_into_tensor, all_reduce, @@ -56,23 +45,16 @@ ) from lightllm.server.core.objs.shm_objs_io_buffer import ShmObjsIOBuffer from lightllm.server.router.model_infer.mode_backend.overlap_events import OverlapEventManager, OverlapEventPack -from lightllm.models.deepseek_mtp.model import Deepseek3MTPModel -from lightllm.models.qwen3_moe_mtp.model import Qwen3MOEMTPModel -from lightllm.models.mistral_mtp.model import MistralMTPModel -from lightllm.models.glm4_moe_lite_mtp.model import Glm4MoeLiteMTPModel from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample from lightllm.common.basemodel.triton_kernel.gather_token_id import scatter_token from lightllm.server.pd_io_struct import PDChunckedTransTaskRet from .multi_level_kv_cache import MultiLevelKvCacheModule from lightllm.utils.profiler import ProcessProfiler, ProfilerCmd -logger = init_logger(__name__) - class ModeBackend: def __init__(self) -> None: self.shm_req_manager = ShmReqManager() - start_args = get_env_start_args() self.overlap_event_manager = OverlapEventManager() # 标识是否支持 overlap 功能,很多子类模式如 xgrammar 和 outlines 当前不支持 overlap 高性能模式 @@ -85,11 +67,9 @@ def __init__(self) -> None: # extra_post_req_handle_func 用于添加请求InferReq的状态变化中添加额外的后处理信息,主要是状态机相关的调整等。 self.extra_post_req_handle_func: Optional[Callable[[InferReq, int, float], None]] = None - self.enable_decode_microbatch_overlap = start_args.enable_decode_microbatch_overlap - self.enable_prefill_microbatch_overlap = start_args.enable_prefill_microbatch_overlap - self.spec_config = SpeculativeConfig.from_args(start_args, dynamic_verify=enable_dynamic_mtp_verify()) - self.spec_config.validate() - self.spec_adapter = None + self.enable_decode_microbatch_overlap = get_env_start_args().enable_decode_microbatch_overlap + self.enable_prefill_microbatch_overlap = get_env_start_args().enable_prefill_microbatch_overlap + self.spec_engine = None # 控制 _get_classed_reqs 分类的参数变量,不同的 backend 具有可能需要不同的分类运行条件。 self.classed_req_no_decode = False @@ -104,7 +84,6 @@ def __init__(self) -> None: self._radix_tree_merge_update_delta: int = get_radix_tree_merge_update_delta() pass - def init_model(self, kvargs): self.args: StartArgs = kvargs.get("args", None) assert self.args is not None @@ -128,22 +107,17 @@ def init_model(self, kvargs): self.is_multinode_tp = self.args.nnodes > 1 and self.args.dp == 1 self.is_pd_mode = self.run_mode in ["prefill", "decode"] self.is_pd_decode_mode = self.run_mode == "decode" - self.spec_config = SpeculativeConfig.from_args(self.args, dynamic_verify=enable_dynamic_mtp_verify()) - self.spec_config.validate() - if self.spec_config.needs_target_layer_hidden: + if self.args.mtp_mode in ("eagle3", "dspark", "dflash"): assert ( not self.args.enable_decode_microbatch_overlap - ), f"{self.spec_config.mode} mode does not support decode microbatch overlap" + ), f"{self.args.mtp_mode} mode does not support decode microbatch overlap" assert ( not self.args.enable_prefill_microbatch_overlap - ), f"{self.spec_config.mode} mode does not support prefill microbatch overlap" + ), f"{self.args.mtp_mode} mode does not support prefill microbatch overlap" self.logger = init_logger(__name__) self.weight_dir = kvargs["weight_dir"] - self._normalize_block_mtp_step_from_first_draft_config() - # p d 分离模式,decode节点才会使用的参数 - self.pd_rpyc_ports = kvargs.get("pd_rpyc_ports", None) max_total_token_num = kvargs["max_total_token_num"] init_distributed_env(kvargs) @@ -160,6 +134,7 @@ def init_model(self, kvargs): model_cfg, _ = PretrainedConfig.get_config_dict(self.weight_dir) + target_decode_batch_multiplier = self.args.mtp_step + 1 if self.args.mtp_mode is not None else 1 model_kvargs = { "weight_dir": self.weight_dir, "max_total_token_num": max_total_token_num, @@ -171,8 +146,6 @@ def init_model(self, kvargs): "disable_chunked_prefill": self.disable_chunked_prefill, "data_type": kvargs.get("data_type", "float16"), "graph_max_batch_size": kvargs.get("graph_max_batch_size", 16), - "graph_split_batch_size": kvargs.get("graph_split_batch_size", self.args.graph_split_batch_size), - "graph_grow_step_size": kvargs.get("graph_grow_step_size", self.args.graph_grow_step_size), "graph_max_len_in_batch": kvargs.get("graph_max_len_in_batch", 8196), "disable_cudagraph": kvargs.get("disable_cudagraph", False), "mem_fraction": kvargs.get("mem_fraction", 0.9), @@ -181,12 +154,13 @@ def init_model(self, kvargs): "quant_cfg": kvargs.get("quant_cfg", None), "expert_dtype": kvargs.get("expert_dtype", None), "run_mode": self.run_mode, + "hidden_layer_ids": self._target_hidden_layer_ids(model_cfg), + "decode_batch_multiplier": target_decode_batch_multiplier, } self.model, self.is_multimodal = get_model(model_cfg, model_kvargs) self.model: TpPartBaseModel = self.model # for easy typing set_random_seed(2147483647) self.is_linear_att_mixed_model = isinstance(self.model.req_manager, ReqManagerForMamba) - self._validate_linear_att_spec_support() if self.is_linear_att_mixed_model: self.linear_att_cache_manager = LinearAttCacheManager( @@ -253,8 +227,8 @@ def init_model(self, kvargs): ) if self.args.run_mode in ["prefill", "decode"] or self.args.enable_dp_prompt_cache_fetch: - # The target manager already includes the speculative full-attention - # layer slots, so it can be shared before draft model initialization. + # 如果存在需要跨进程使用mem manger的特性,则将mem manager写入到 shm中,方便 + # 读取 self.model.mem_manager.write_to_shm(req_manager=self.model.req_manager) dist.barrier(group=self.node_nccl_group) @@ -281,12 +255,9 @@ def init_model(self, kvargs): self.shm_pd_trans_io_buffer = ShmObjsIOBuffer(tail_str="pd") # 开启 mtp 模式,需要完成mtp model的初始化 - if self.spec_config.enabled: - self.init_mtp_draft_model(kvargs) - if self.spec_config.dynamic_verify: - g_infer_context.init_dynamic_mtp_planner(mtp_step=self.mtp_step, mode=self.spec_config.mode) - self.spec_adapter = build_spec_runtime(self) - self._attach_spec_adapter() + if self.args.mtp_mode is not None: + self.init_spec_draft_model(model_kvargs) + self.spec_engine = SpecEngine(backend=self) if self.args.enable_cpu_cache: self.multi_level_cache_module = MultiLevelKvCacheModule(self) @@ -339,45 +310,31 @@ def prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): def decode(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): raise NotImplementedError() - def init_mtp_draft_model(self, main_kvargs: dict): - self.mtp_step = self.args.mtp_step + def init_spec_draft_model(self, main_kvargs: dict): + self.max_draft_step = self.args.mtp_step self.draft_models = [] - spec_config = self.spec_config + spec_mode = self.args.mtp_mode + is_chained_draft = spec_mode in ("vanilla_with_att", "vanilla_no_att", "qwen3next_vanilla") + is_recurrent_draft = spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3", "qwen3next_eagle") os.environ["DISABLE_CHECK_MAX_LEN_INFER"] = "1" - num_mtp_modules = spec_config.draft_model_count - mtp_draft_model_dirs = self.args.mtp_draft_model_dir - if isinstance(mtp_draft_model_dirs, str): - mtp_draft_model_dirs = [mtp_draft_model_dirs] - assert mtp_draft_model_dirs is not None - assert len(mtp_draft_model_dirs) >= num_mtp_modules - - draft_graph_max_override = getattr(self.args, "mtp_draft_graph_max_batch_size", None) - draft_graph_split_override = getattr(self.args, "mtp_draft_graph_split_batch_size", None) - draft_graph_grow_override = getattr(self.args, "mtp_draft_graph_grow_step_size", None) - draft_graph_max_batch_size = ( - draft_graph_max_override - if draft_graph_max_override is not None - else main_kvargs.get("graph_max_batch_size", 16) - ) - draft_graph_split_batch_size = ( - draft_graph_split_override - if draft_graph_split_override is not None - else main_kvargs.get("graph_split_batch_size", self.args.graph_split_batch_size) - ) - draft_graph_grow_step_size = ( - draft_graph_grow_override - if draft_graph_grow_override is not None - else main_kvargs.get("graph_grow_step_size", self.args.graph_grow_step_size) - ) - - for i in range(num_mtp_modules): - mtp_model_cfg, _ = PretrainedConfig.get_config_dict(mtp_draft_model_dirs[i]) - spec_config = self.spec_config - model_type = mtp_model_cfg.get("model_type", "") - mtp_model_kvargs = { - "weight_dir": mtp_draft_model_dirs[i], + draft_model_count = self.max_draft_step if is_chained_draft else 1 + draft_model_dirs = self.args.mtp_draft_model_dir + assert draft_model_dirs is not None + assert len(draft_model_dirs) >= draft_model_count + + for i in range(draft_model_count): + draft_model_cfg, _ = PretrainedConfig.get_config_dict(draft_model_dirs[i]) + if is_chained_draft: + draft_decode_batch_multiplier = self.max_draft_step + 1 + elif is_recurrent_draft: + draft_decode_batch_multiplier = 1 + else: + block_size = int(draft_model_cfg["block_size"]) + draft_decode_batch_multiplier = block_size + draft_model_kvargs = { + "weight_dir": draft_model_dirs[i], "max_total_token_num": self.model.mem_manager.size, "load_way": main_kvargs["load_way"], "max_req_num": main_kvargs.get("max_req_num", 1000), @@ -386,9 +343,7 @@ def init_mtp_draft_model(self, main_kvargs: dict): "return_all_prompt_logics": False, "disable_chunked_prefill": self.disable_chunked_prefill, "data_type": main_kvargs.get("data_type", "float16"), - "graph_max_batch_size": draft_graph_max_batch_size, - "graph_split_batch_size": draft_graph_split_batch_size, - "graph_grow_step_size": draft_graph_grow_step_size, + "graph_max_batch_size": main_kvargs["graph_max_batch_size"], "graph_max_len_in_batch": main_kvargs.get("graph_max_len_in_batch", 8196), "disable_cudagraph": main_kvargs.get("disable_cudagraph", False), "mem_fraction": main_kvargs["mem_fraction"], @@ -399,125 +354,17 @@ def init_mtp_draft_model(self, main_kvargs: dict): "run_mode": "normal", "main_model": self.model, "mtp_previous_draft_models": self.draft_models.copy(), + "decode_batch_multiplier": draft_decode_batch_multiplier, } - model_type = mtp_model_cfg.get("model_type", "") - if model_type == "deepseek_v3": - assert spec_config.uses_attention_draft - self.draft_models.append(Deepseek3MTPModel(mtp_model_kvargs)) - elif model_type == "qwen3_moe": - assert spec_config.uses_no_attention_draft and not spec_config.is_eagle3 - self.draft_models.append(Qwen3MOEMTPModel(mtp_model_kvargs)) - elif model_type == "mistral": - assert spec_config.uses_no_attention_draft and not spec_config.is_eagle3 - self.draft_models.append(MistralMTPModel(mtp_model_kvargs)) - elif model_type == "glm4_moe_lite": - assert spec_config.uses_attention_draft - self.draft_models.append(Glm4MoeLiteMTPModel(mtp_model_kvargs)) - elif model_type in ("qwen3_5", "qwen3_5_text"): - assert spec_config.uses_attention_draft - from lightllm.models.qwen3_5_mtp.model import Qwen3_5MTPModel - - self.draft_models.append(Qwen3_5MTPModel(mtp_model_kvargs)) - elif model_type in ("qwen3_5_moe", "qwen3_5_moe_text"): - assert spec_config.uses_attention_draft - from lightllm.models.qwen3_5_moe_mtp.model import Qwen3_5MoeMTPModel - - self.draft_models.append(Qwen3_5MoeMTPModel(mtp_model_kvargs)) - elif spec_config.is_eagle3 and is_eagle3_draft_config(mtp_model_cfg): - from lightllm.models.qwen3_eagle.model import Qwen3EagleModel - - self.draft_models.append(Qwen3EagleModel(mtp_model_kvargs)) - elif spec_config.is_dflash and is_qwen3_5_dflash_draft_config(mtp_model_cfg): - from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel - - self.draft_models.append(Qwen3_5DFlashModel(mtp_model_kvargs)) - elif spec_config.is_dflash and is_qwen3_dflash_draft_config(mtp_model_cfg): - from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel - - self.draft_models.append(Qwen3DFlashModel(mtp_model_kvargs)) - elif spec_config.is_dspark and is_qwen3_dspark_draft_config(mtp_model_cfg): - if self.is_linear_att_mixed_model: - from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel - - self.draft_models.append(Qwen3_5DSparkModel(mtp_model_kvargs)) - else: - from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel - - self.draft_models.append(Qwen3DSparkModel(mtp_model_kvargs)) - elif (spec_config.is_dflash or spec_config.is_dspark) and is_gemma4_dspark_draft_config(mtp_model_cfg): - raise NotImplementedError("Gemma4 DSpark draft checkpoints are not wired to LightLLM serving yet.") - elif (spec_config.is_dflash or spec_config.is_dspark) and is_dspark_draft_config(mtp_model_cfg): - raise ValueError(f"Unsupported DSpark-family draft architecture: {mtp_model_cfg.get('architectures')}") - else: - raise ValueError(f"Unsupported MTP model type: {model_type}") - - self.logger.info(f"loaded mtp model class {self.draft_models[i].__class__}") - return - - def _normalize_block_mtp_step_from_config(self, mtp_model_cfg: dict) -> None: - if not self.spec_config.uses_block_draft_model: - return - - layout = get_block_draft_layout( - mtp_model_cfg, - mode=self.spec_config.mode, - require_confidence_head=self.spec_config.is_dspark, - ) - self.block_draft_layout = layout - configured_step = int(getattr(self.args, "mtp_step", 0)) - draft_step = layout.resolve_draft_step(configured_step) - if configured_step not in (0, draft_step): - self.logger.warning( - "Overriding mtp_step=%s with draft_step=%s from block_size=%s for %s mode", - configured_step, - draft_step, - layout.query_block_size, - self.spec_config.mode, + draft_model_class = get_draft_model_class( + model_cfg=draft_model_cfg, + spec_mode=spec_mode, + is_linear_att_mixed_model=self.is_linear_att_mixed_model, ) - self.args.mtp_step = draft_step - self.mtp_step = draft_step - self.spec_config = replace(self.spec_config, step=draft_step) - return - - def _validate_linear_att_spec_support(self) -> None: - """Validate block-draft checkpoint families against hybrid targets.""" - - if not self.spec_config.enabled or not self.spec_config.uses_block_draft_model: - return + self.draft_models.append(draft_model_class(draft_model_kvargs)) - mtp_draft_model_dirs = self.args.mtp_draft_model_dir - if isinstance(mtp_draft_model_dirs, str): - mtp_draft_model_dirs = [mtp_draft_model_dirs] - assert mtp_draft_model_dirs is not None and len(mtp_draft_model_dirs) == 1 - mtp_model_cfg, _ = PretrainedConfig.get_config_dict(mtp_draft_model_dirs[0]) - is_qwen35_dflash = is_qwen3_5_dflash_draft_config(mtp_model_cfg) - is_qwen_dspark = is_qwen3_dspark_draft_config(mtp_model_cfg) - if self.is_linear_att_mixed_model: - if self.spec_config.is_dflash: - assert is_qwen35_dflash, ( - "linear-attention mixed targets require a Qwen3_5DFlashModel checkpoint in DFlash mode, " - f"got architectures={mtp_model_cfg.get('architectures')}" - ) - else: - assert is_qwen_dspark, ( - "linear-attention mixed targets require a Qwen3DSparkModel checkpoint in DSpark mode, " - f"got architectures={mtp_model_cfg.get('architectures')}" - ) - else: - assert not is_qwen35_dflash, "Qwen3_5DFlashModel requires a Qwen3Next target" - return - - def _normalize_block_mtp_step_from_first_draft_config(self) -> None: - if not self.spec_config.uses_block_draft_model: - return - - mtp_draft_model_dirs = self.args.mtp_draft_model_dir - if isinstance(mtp_draft_model_dirs, str): - mtp_draft_model_dirs = [mtp_draft_model_dirs] - assert mtp_draft_model_dirs is not None and len(mtp_draft_model_dirs) > 0 - mtp_model_cfg, _ = PretrainedConfig.get_config_dict(mtp_draft_model_dirs[0]) - self._normalize_block_mtp_step_from_config(mtp_model_cfg) + self.logger.info(f"loaded speculative draft model class {self.draft_models[i].__class__}") return def _async_copy_next_token_infos_to_pin_mem( @@ -622,14 +469,26 @@ def _capture_prompt_logprobs_if_needed( start_loc += q_len return - def _attach_spec_adapter(self) -> None: - if not self.spec_config.enabled: - return - assert self.spec_adapter is not None - self.model.set_spec_adapter(self.spec_adapter) - for draft_model in self.draft_models: - draft_model.set_spec_adapter(self.spec_adapter) - return + def _target_hidden_layer_ids(self, model_cfg: dict): + if self.args.mtp_mode not in ("eagle3", "dspark", "dflash"): + return None + + draft_model_dirs = self.args.mtp_draft_model_dir + assert draft_model_dirs + draft_cfg, _ = PretrainedConfig.get_config_dict(draft_model_dirs[0]) + target_layer_ids = draft_cfg.get("target_layer_ids") + if target_layer_ids is None and self.args.mtp_mode == "dflash": + target_layer_ids = draft_cfg.get("dflash_config", {}).get("target_layer_ids") + if target_layer_ids is None: + layer_num = int(model_cfg.get("num_hidden_layers", model_cfg.get("n_layer"))) + target_layer_ids = [1, layer_num // 2 - 1, layer_num - 4] + + layer_num = int(model_cfg.get("num_hidden_layers", model_cfg.get("n_layer"))) + target_layer_ids = tuple(int(layer_id) for layer_id in target_layer_ids) + assert target_layer_ids and all( + 0 <= layer_id < layer_num for layer_id in target_layer_ids + ), f"invalid target_layer_ids={target_layer_ids} for target layer_num={layer_num}" + return target_layer_ids def _try_read_new_reqs(self): if self.is_multinode_tp: @@ -984,7 +843,6 @@ def _post_handle( run_reqs_update_packs: List[InferReqUpdatePack], extra_post_req_handle_func: Optional[Callable[[InferReq, int, float], None]] = None, pd_prefill_chunked_handle_func: Optional[Callable[[InferReq, int, float, int], None]] = None, - count_mtp_accepted_tokens: bool = False, ): """ extra_post_req_handle_func 用于提供在一个请求确定输出的时候,给出额外的后处理操作,主要是用于 @@ -997,17 +855,11 @@ def _post_handle( if isinstance(next_token_ranks, torch.Tensor): next_token_ranks = next_token_ranks.numpy() - mtp_seen_count = collections.Counter() if count_mtp_accepted_tokens and self.is_master_in_dp else None for req_obj, next_token_id, next_token_logprob, next_token_rank, pack in zip( run_reqs, next_token_ids, next_token_logprobs, next_token_ranks, run_reqs_update_packs ): req_obj: InferReq = req_obj pack: InferReqUpdatePack = pack - if mtp_seen_count is not None: - seen_count = mtp_seen_count[req_obj.req_idx] - mtp_seen_count[req_obj.req_idx] += 1 - if seen_count > 0 and not req_obj.finish_status.is_finished(): - req_obj.update_mtp_accepted_token_num(accept_token_num=1) pack.handle( next_token_id=next_token_id, next_token_logprob=next_token_logprob, @@ -1033,45 +885,42 @@ def _filter_reqs(self, reqs: List[InferReq]): def _trans_req_ids_to_req_objs(self, req_ids: List[int]) -> List[InferReq]: return [g_infer_context.requests_mapping[req_id] for req_id in req_ids] - def _update_mtp_accept_ratio( + def _update_spec_accept_ratio( self, decode_reqs: List[InferReq], - mtp_accept_len_cpu: torch.Tensor, + spec_accept_len_cpu: torch.Tensor, ): if self.is_master_in_dp: - for req, accept_len in zip(decode_reqs, mtp_accept_len_cpu.numpy()): - req.update_mtp_accepted_token_num(accept_token_num=max(int(accept_len) - 1, 0)) + for req, accept_len in zip(decode_reqs, spec_accept_len_cpu): + req.update_spec_accepted_token_num(accept_token_num=accept_len - 1) return - def _update_mtp_verify_token_num( - self, decode_reqs: List[InferReq], dynamic_mtp_run_reqs: Optional[List[InferReq]] = None + def _update_spec_verify_token_num( + self, decode_reqs: List[InferReq], selected_run_reqs: Optional[List[InferReq]] = None ): - if self.is_master_in_dp: - if dynamic_mtp_run_reqs is None: - for req in decode_reqs: - assert req.mtp_step > 0 - verify_len = 1 + req.mtp_step - req.update_mtp_verify_token_num(verify_token_num=verify_len) - req.update_mtp_verify_step_num(verify_step_num=1) - else: - counter = collections.Counter([req.req_idx for req in dynamic_mtp_run_reqs]) - for req in decode_reqs: - verify_token_num = counter[req.req_idx] - if verify_token_num > 0: - req.update_mtp_verify_token_num(verify_token_num=verify_token_num) - req.update_mtp_verify_step_num(verify_step_num=1) - return + if not self.is_master_in_dp: + return + + if selected_run_reqs is None: + for req in decode_reqs: + req.update_spec_verify_token_num(verify_token_num=req.max_draft_step + 1) + req.update_spec_verify_step_num(verify_step_num=1) + return + + verify_rows_by_req = Counter(req.req_idx for req in selected_run_reqs) + for req in decode_reqs: + verify_token_num = verify_rows_by_req[req.req_idx] + if verify_token_num > 0: + req.update_spec_verify_token_num(verify_token_num=verify_token_num) + req.update_spec_verify_step_num(verify_step_num=1) def _gen_argmax_token_ids(self, model_output: ModelOutput): - if model_output.mtp_draft_token_ids is not None: - draft_next_token_ids_gpu = model_output.mtp_draft_token_ids - else: - logits = model_output.logits - draft_next_token_ids_gpu = torch.argmax(logits, dim=-1) + logits = model_output.logits + draft_next_token_ids_gpu = torch.argmax(logits, dim=-1) # 如果draft和target的词表不同,需要把draft token映射回主模型词表。 - if self.spec_config.needs_draft_vocab_mapping: + if self.args.mtp_mode == "eagle3": draft_next_token_ids_gpu = self.draft_models[0].map_draft_vocab_to_main_vocab(draft_next_token_ids_gpu) return draft_next_token_ids_gpu @@ -1081,7 +930,7 @@ def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): max_probs, draft_next_token_ids_gpu = torch.max(probs, dim=-1) # 如果self.d2t不为None,那么draft的token需要进行相应的转换 - if self.spec_config.needs_draft_vocab_mapping: + if self.args.mtp_mode == "eagle3": draft_next_token_ids_gpu = self.draft_models[0].map_draft_vocab_to_main_vocab(draft_next_token_ids_gpu) return draft_next_token_ids_gpu, max_probs diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index c219d55ee9..cbffdbde9a 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -15,6 +15,7 @@ from lightllm.utils.dist_utils import get_current_device_id from .control_state import ControlState from lightllm.utils.dist_utils import create_new_group_for_current_dp +from lightllm.utils.envs_utils import enable_dynamic_spec, get_env_start_args logger = init_logger(__name__) @@ -25,13 +26,13 @@ def __init__(self) -> None: # 用于控制每一步是执行prefill 和 decode 还是跳过 self.control_state_machine = ControlState() - self.enable_dynamic_mtp = False + self.enable_dynamic_spec = False # 在 mtp 模式下切换绑定的prefill 和 decode 函数 - if self.spec_config.enabled: - self.prefill = self.prefill_mtp - self.decode = self.decode_mtp - self.enable_dynamic_mtp = self.spec_config.dynamic_verify + if get_env_start_args().mtp_mode is not None: + self.prefill = self.prefill_spec + self.decode = self.decode_spec + self.enable_dynamic_spec = enable_dynamic_spec() else: self.prefill = self.prefill_normal self.decode = self.decode_normal @@ -39,12 +40,15 @@ def __init__(self) -> None: self.classed_req_strict_prefill = False return + # cpu 把算子提交到gpu 上 + # GPU + # CPU + def init_custom(self): super().init_custom() - if self.enable_dynamic_mtp: - self.mtp_gloo_group = create_new_group_for_current_dp("gloo") - logger.info(f"mtp_gloo_group ranks {dist.get_rank(self.mtp_gloo_group)}") - return + if self.enable_dynamic_spec: + self.spec_gloo_group = create_new_group_for_current_dp("gloo") + logger.info(f"spec_gloo_group ranks {dist.get_rank(self.spec_gloo_group)}") def infer_loop(self): torch.cuda.set_device(get_current_device_id()) @@ -181,7 +185,7 @@ def decode_normal( event_pack.notify_pre_post_handle() return - def prefill_mtp( + def prefill_spec( self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq], @@ -205,9 +209,10 @@ def prefill_mtp( mask_func=self.prefill_mask_func, ) # mtp kv fill - spec_runtime = self.spec_adapter - spec_runtime.build_initial_draft_state( + spec_engine = self.spec_engine + spec_engine.build_initial_draft_state( model_input=model_input, + model_output=model_output, next_token_ids=next_token_ids, ) g_infer_context.copy_linear_att_state_to_cache_buffer( @@ -239,26 +244,24 @@ def prefill_mtp( event_pack.notify_pre_post_handle() return - def decode_mtp( + def decode_spec( self, event_pack: OverlapEventPack, decode_reqs: List[InferReq], ): - """ - MTP解码的通用流程,整合eagle和vanilla的共同逻辑 - """ + """Run the shared speculative draft-and-verify decode flow.""" model_input, run_reqs = prepare_decode_inputs(decode_reqs) - spec_runtime = self.spec_adapter + spec_engine = self.spec_engine with torch.cuda.stream(g_infer_context.get_overlap_stream()): - spec_plan = spec_runtime.plan_decode(model_input=model_input, req_num=len(decode_reqs)) + spec_plan = spec_engine.plan_decode(model_input=model_input, req_num=len(decode_reqs)) - model_input, selected_run_reqs = spec_runtime.prepare_decode_model_input( + model_input, selected_row_mask = spec_engine.prepare_decode_model_input( model_input=model_input, req_num=len(decode_reqs), plan=spec_plan, ) - selected_run_reqs_cpu = spec_runtime.async_copy_selected_run_reqs(selected_run_reqs) + selected_row_mask_cpu = spec_engine.async_copy_selected_row_mask(selected_row_mask) model_output = self.model.forward(model_input) @@ -266,18 +269,17 @@ def decode_mtp( model_output.logits, run_reqs, self.eos_id, - dynamic_batch_size=spec_plan.dynamic_batch_size, - selected_run_reqs=selected_run_reqs, + selected_row_mask=selected_row_mask, ) next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids) - spec_decode_state = spec_runtime.run_decode_speculative_forward( + spec_decode_state = spec_engine.run_decode_speculative_forward( model_input=model_input, model_output=model_output, run_reqs=run_reqs, req_num=len(decode_reqs), plan=spec_plan, - selected_run_reqs_cpu=selected_run_reqs_cpu, + selected_row_mask_cpu=selected_row_mask_cpu, next_token_ids=next_token_ids, next_token_logprobs=next_token_logprobs, next_token_ranks=next_token_ranks, @@ -287,26 +289,26 @@ def decode_mtp( # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() - run_reqs, verify_ok_reqs = spec_runtime.resolve_decode_pre_post_reqs( + run_reqs, verify_ok_reqs = spec_engine.resolve_decode_pre_post_reqs( state=spec_decode_state, decode_reqs=decode_reqs, ) - self._update_mtp_verify_token_num( + self._update_spec_verify_token_num( decode_reqs=decode_reqs, - dynamic_mtp_run_reqs=run_reqs if self.enable_dynamic_mtp else None, + selected_run_reqs=run_reqs if self.enable_dynamic_spec else None, ) update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) # 第三阶段 event_pack.notify_forward_and_wait_post_handle() - spec_post_state = spec_runtime.finish_decode_post( + spec_post_state = spec_engine.finish_decode_post( state=spec_decode_state, req_num=len(decode_reqs), run_reqs=run_reqs, ) - self._update_mtp_accept_ratio( + self._update_spec_accept_ratio( decode_reqs=decode_reqs, - mtp_accept_len_cpu=spec_post_state.mtp_accept_len_cpu, + spec_accept_len_cpu=spec_post_state.spec_accept_len_cpu, ) self._post_handle( diff --git a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py index 5bec2db491..8c4dfc1a33 100644 --- a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py @@ -1,4 +1,5 @@ import torch +from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.infer_batch import ( g_infer_context, InferReq, @@ -13,16 +14,22 @@ from lightllm.common.basemodel.triton_kernel.gather_token_id import scatter_token from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from ..chunked_prefill.impl import ChunkedPrefillBackend +from lightllm.utils.envs_utils import get_env_start_args class DiversehBackend(ChunkedPrefillBackend): def __init__(self) -> None: super().__init__() - if self.spec_config.enabled: - # 当前只有 mistral mtp 可以使用 diverse mode 的 mtp 功能。 + if get_env_start_args().mtp_mode: + # 当前只有 mistral 和 Qwen3Next mtp 可以使用 diverse mode 的 mtp 功能。 self.prefill = self.beam_prefill - assert self.spec_config.uses_no_attention_draft and not self.spec_config.is_eagle3 + assert get_env_start_args().mtp_mode in [ + "vanilla_no_att", + "eagle_no_att", + "qwen3next_vanilla", + "qwen3next_eagle", + ] else: self.prefill = self.beam_prefill diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 4ab9c7e7d0..be87b8725b 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -13,9 +13,10 @@ ) from lightllm.server.router.model_infer.mode_backend.overlap_events import OverlapEventPack from lightllm.utils.dist_utils import get_current_device_id +from lightllm.utils.envs_utils import get_env_start_args from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.common.basemodel.triton_kernel.mtp_utils import ( - linear_att_mtp_state_index_update, + linear_att_spec_state_index_update, ) from .control_state import DPControlState @@ -28,23 +29,32 @@ def __init__(self) -> None: self.control_state_machine = DPControlState(backend=self) # 在 mtp 模式下切换绑定的prefill 和 decode 函数 - if self.spec_config.enabled: - if self.spec_config.uses_block_draft_model: + spec_mode = get_env_start_args().mtp_mode + if spec_mode is not None: + if spec_mode in ("dspark", "dflash"): raise NotImplementedError("DP backend does not support DFlash/DSpark block draft mode yet.") - self.is_mtp_eagle = self.spec_config.uses_recurrent_draft_model - self.num_mtp_models = self.spec_config.draft_model_count + self.uses_recurrent_draft = spec_mode in ( + "eagle_with_att", + "eagle_no_att", + "eagle3", + "qwen3next_eagle", + ) if self.enable_prefill_microbatch_overlap: - self.prefill = self.prefill_overlap_mtp + self.prefill = self.prefill_overlap_spec else: - self.prefill = self.prefill_mtp + self.prefill = self.prefill_spec if self.enable_decode_microbatch_overlap: - self.decode = self.decode_overlap_mtp + self.decode = self.decode_overlap_spec self._draft_decode_overlap_func = ( - self._draft_decode_eagle_overlap if self.is_mtp_eagle else self._draft_decode_vanilla_overlap + self._draft_decode_eagle_overlap + if self.uses_recurrent_draft + else self._draft_decode_vanilla_overlap ) else: - self.decode = self.decode_mtp - self._draft_decode_func = self._draft_decode_eagle if self.is_mtp_eagle else self._draft_decode_vanilla + self.decode = self.decode_spec + self._draft_decode_func = ( + self._draft_decode_eagle if self.uses_recurrent_draft else self._draft_decode_vanilla + ) else: if self.enable_prefill_microbatch_overlap: self.prefill = self.prefill_overlap @@ -59,6 +69,28 @@ def __init__(self) -> None: self.classed_req_strict_prefill = False return + @staticmethod + def _build_padded_next_token_ids( + token_ids: torch.Tensor, + batch_size: int, + copy_len: int, + source_start: int = 0, + ) -> torch.Tensor: + """Pad DP-local draft tokens to the collective batch shape.""" + + copy_len = int(copy_len) + source_start = int(source_start) + assert 0 <= copy_len <= int(batch_size) + assert source_start + copy_len <= token_ids.shape[0] + + padded_token_ids = torch.zeros((int(batch_size),), dtype=torch.int64, device=token_ids.device) + if copy_len > 0: + padded_token_ids[:copy_len].copy_( + token_ids[source_start : source_start + copy_len], + non_blocking=True, + ) + return padded_token_ids + def _init_reqs(self, reqs: List[Tuple]): if not self.args.enable_dp_prompt_cache_fetch: return super()._init_reqs(reqs) @@ -398,7 +430,7 @@ def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferRe event_pack.notify_pre_post_handle() return - def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): + def prefill_spec(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): # main model prefill model_input, run_reqs, _ = padded_prepare_prefill_inputs(prefill_reqs) req_num = len(run_reqs) @@ -426,14 +458,14 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] ) # mtp kv fill - draft_next_token_ids_gpu = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num > 0 else None, + draft_next_token_ids_gpu = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input.batch_size, copy_len=req_num, - device=model_input.b_req_idx.device, ) - self.spec_adapter.build_initial_draft_state( + self.spec_engine.build_initial_draft_state( model_input=model_input, + model_output=model_output, next_token_ids=draft_next_token_ids_gpu, ) if req_num > 0: @@ -470,14 +502,14 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] event_pack.notify_pre_post_handle() return - def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): + def decode_spec(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): model_input, run_reqs, _ = padded_prepare_decode_inputs(decode_reqs) b_mtp_index_cpu = model_input.b_mtp_index req_num = len(run_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output = self.model.forward(model_input) - mtp_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None + spec_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None if req_num > 0: logits = model_output.logits[0:req_num, :] b_mtp_index_cpu = b_mtp_index_cpu[0:req_num] @@ -499,29 +531,29 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): dtype=torch.int32, ).cuda(non_blocking=True) - verify_result = self.spec_adapter.verify_target_tokens( + verify_result = self.spec_engine.verify_target_tokens( new_next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, ) - mtp_accept_len = verify_result.accept_len + spec_accept_len = verify_result.accept_len accepted_index = verify_result.accepted_index if self.is_linear_att_mixed_model: - linear_att_mtp_state_index_update( + linear_att_spec_state_index_update( req_to_mtp_state_index=self.model.req_manager.req_to_mtp_state_index, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, b_mtp_index=model_input.b_mtp_index[0:req_num], accepted_index=accepted_index, - max_mtp_step=self.mtp_step + 1, + verify_width=self.max_draft_step + 1, ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", gpu_tensor=accepted_index, ) - mtp_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="mtp_accept_len", - gpu_tensor=mtp_accept_len, + spec_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="spec_accept_len", + gpu_tensor=spec_accept_len, ) verify_event = torch.cuda.Event() @@ -529,9 +561,10 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): eagle_mem_indexes_cpu = self._draft_decode_func( model_input=model_input, + model_output=model_output, next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - mtp_accept_len=mtp_accept_len, + spec_accept_len=spec_accept_len, req_num=req_num, ) if req_num > 0: @@ -548,18 +581,18 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() - self._update_mtp_verify_token_num(decode_reqs=decode_reqs) + self._update_spec_verify_token_num(decode_reqs=decode_reqs) verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) # 第三阶段 event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() - self._update_mtp_accept_ratio(decode_reqs=decode_reqs, mtp_accept_len_cpu=mtp_accept_len_cpu) need_free_mem_indexes = model_input.mem_indexes_cpu[0:req_num][accepted_index_cpu == 0] if eagle_mem_indexes_cpu is not None: need_free_mem_indexes = torch.cat([need_free_mem_indexes, eagle_mem_indexes_cpu], dim=0) + self._update_spec_accept_ratio(decode_reqs=decode_reqs, spec_accept_len_cpu=spec_accept_len_cpu) select_mask = torch.tensor(accepted_index_cpu, dtype=torch.bool, device="cpu") self._post_handle( run_reqs=verify_ok_reqs, @@ -583,41 +616,46 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): def _draft_decode_vanilla( self, model_input: ModelInput, + model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, - mtp_accept_len: torch.Tensor, + spec_accept_len: torch.Tensor, req_num: int, ): all_next_token_ids = [] # share some inference info with the main model draft_model_input = model_input - draft_next_token_ids_gpu = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num > 0 else None, + draft_hidden = model_output.spec_hidden + assert draft_hidden is not None + draft_next_token_ids_gpu = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input.batch_size, copy_len=req_num, - device=model_input.b_req_idx.device, ) all_next_token_ids.append(draft_next_token_ids_gpu) # process the draft model output - for draft_model_idx in range(self.mtp_step): + for draft_model_idx in range(self.max_draft_step): - draft_model_input = self.spec_adapter.prepare_draft_decode_input( + draft_model_input = self.spec_engine.prepare_draft_decode_input( model_input=draft_model_input, next_token_ids=draft_next_token_ids_gpu, + mtp_draft_input_hiddens=draft_hidden, ) # spec decode: MTP draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) + draft_hidden = draft_model_output.spec_hidden + assert draft_hidden is not None draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) all_next_token_ids.append(draft_next_token_ids_gpu) if req_num > 0: - self.spec_adapter.scatter_token_id_steps( + self.spec_engine.scatter_token_id_steps( token_id_steps=all_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx[:req_num], - mtp_accept_len=mtp_accept_len, + spec_accept_len=spec_accept_len, row_count=req_num, ) return None @@ -625,62 +663,61 @@ def _draft_decode_vanilla( def _draft_decode_eagle( self, model_input: ModelInput, + model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, - mtp_accept_len: torch.Tensor, + spec_accept_len: torch.Tensor, req_num: int, ): - all_next_token_ids = [] - # share some inference info with the main model - draft_model_input = model_input - all_next_token_ids.append(next_token_ids) - draft_next_token_ids_gpu = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num > 0 else None, + verify_width = self.max_draft_step + 1 + assert req_num % verify_width == 0 + assert model_input.batch_size % verify_width == 0 + real_request_num = req_num // verify_width + request_capacity = model_input.batch_size // verify_width + + padded_next_token_ids = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input.batch_size, copy_len=req_num, + ) + padded_start_locs = torch.arange( + 0, + model_input.batch_size, + verify_width, + dtype=torch.int32, device=model_input.b_req_idx.device, ) - - real_req_num = req_num // (self.mtp_step + 1) - padded_req_num = model_input.batch_size // (self.mtp_step + 1) - real_req_num - eagle_mem_indexes_cpu = self.spec_adapter.alloc_extra_mem_indexes(real_req_num * self.mtp_step) - eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) - - # process the draft model output - for _step in range(self.mtp_step): - - draft_model_input = self.spec_adapter.prepare_draft_decode_input( - model_input=draft_model_input, - next_token_ids=draft_next_token_ids_gpu, - ) - # spec decode: MTP - draft_model_idx = _step % self.num_mtp_models - draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) - # update the meta info of the inference - draft_model_input.b_seq_len += 1 - draft_model_input.max_kv_seq_len += 1 - draft_model_input.mem_indexes = self.spec_adapter.append_padded_eagle_step_mem_indexes( - model_input=draft_model_input, - eagle_mem_indexes=eagle_mem_indexes, - step=_step, - real_req_num=real_req_num, - padded_req_num=padded_req_num, - mtp_step=self.mtp_step, - ) - draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) - all_next_token_ids.append(draft_next_token_ids_gpu) + padded_accept_len = torch.ones( + (request_capacity,), + dtype=torch.int32, + device=model_input.b_req_idx.device, + ) + if real_request_num > 0: + padded_accept_len[:real_request_num].copy_(spec_accept_len) + + # DP keeps the target verify layout padded for collective shape + # agreement. The proposer still follows the common topology: one + # full-row extend, followed by recurrent decode over one row per + # (real or HOLD) request. + proposal = self.spec_engine.propose_next( + main_model_input=model_input, + main_model_output=model_output, + next_token_ids=padded_next_token_ids, + b_req_mtp_start_loc=padded_start_locs, + draft_step=self.max_draft_step, + accept_len=padded_accept_len, + ) if req_num > 0: - self.spec_adapter.scatter_token_id_steps( - token_id_steps=all_next_token_ids, + self.spec_engine.scatter_next_tokens( b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=proposal.token_ids[:req_num], b_req_idx=model_input.b_req_idx[:req_num], - mtp_accept_len=mtp_accept_len, - row_count=req_num, + spec_accept_len=spec_accept_len, ) - return eagle_mem_indexes_cpu + return proposal.extra_mem_indexes_cpu - def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): + def prefill_overlap_spec(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): ( model_input0, run_reqs0, @@ -722,25 +759,25 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I b_prefill_has_output_cpu=b_has_out_cpu, ) - draft_next_token_ids_gpu0 = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num0 > 0 else None, + draft_next_token_ids_gpu0 = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input0.batch_size, copy_len=req_num0, source_start=0, - device=model_input0.b_req_idx.device, ) - draft_next_token_ids_gpu1 = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num1 > 0 else None, + draft_next_token_ids_gpu1 = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input1.batch_size, copy_len=req_num1, source_start=req_num0, - device=model_input1.b_req_idx.device, ) - self.spec_adapter.build_initial_draft_state_overlap( + self.spec_engine.build_initial_draft_state_overlap( model_input0=model_input0, + model_output0=model_output0, next_token_ids0=draft_next_token_ids_gpu0, model_input1=model_input1, + model_output1=model_output1, next_token_ids1=draft_next_token_ids_gpu1, ) @@ -774,7 +811,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I event_pack.notify_pre_post_handle() return - def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): + def decode_overlap_spec(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): ( model_input0, run_reqs0, @@ -793,7 +830,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf logits0 = model_output0.logits logits1 = model_output1.logits run_reqs = run_reqs0 + run_reqs1 - b_req_idx, mtp_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None, None + b_req_idx, spec_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None, None if (req_num0 + req_num1) > 0: logits = torch.empty( (req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device @@ -817,20 +854,32 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf dtype=torch.int32, ).cuda(non_blocking=True) - verify_result = self.spec_adapter.verify_target_tokens( + verify_result = self.spec_engine.verify_target_tokens( new_next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, ) - mtp_accept_len = verify_result.accept_len + spec_accept_len = verify_result.accept_len accepted_index = verify_result.accepted_index + if self.is_linear_att_mixed_model: + b_mtp_index = torch.cat( + (model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0 + ) + linear_att_spec_state_index_update( + req_to_mtp_state_index=self.model.req_manager.req_to_mtp_state_index, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_req_idx=b_req_idx, + b_mtp_index=b_mtp_index, + accepted_index=accepted_index, + verify_width=self.max_draft_step + 1, + ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", gpu_tensor=accepted_index, ) - mtp_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="mtp_accept_len", - gpu_tensor=mtp_accept_len, + spec_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="spec_accept_len", + gpu_tensor=spec_accept_len, ) all_next_token_ids.append(next_token_ids) @@ -840,9 +889,11 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf eagle_mem_indexes_cpu = self._draft_decode_overlap_func( model_input0=model_input0, model_input1=model_input1, + model_output0=model_output0, + model_output1=model_output1, b_req_idx=b_req_idx, next_token_ids=next_token_ids, - mtp_accept_len=mtp_accept_len, + spec_accept_len=spec_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, req_num0=req_num0, req_num1=req_num1, @@ -860,13 +911,12 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if req_num0 + req_num1 > 0: event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() - self._update_mtp_verify_token_num(decode_reqs=decode_reqs) + self._update_spec_verify_token_num(decode_reqs=decode_reqs) verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() - self._update_mtp_accept_ratio(decode_reqs=decode_reqs, mtp_accept_len_cpu=mtp_accept_len_cpu) mem_indexes_cpu = torch.cat( (model_input0.mem_indexes_cpu[0:req_num0], model_input1.mem_indexes_cpu[0:req_num1]), dim=0 ) @@ -874,6 +924,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if eagle_mem_indexes_cpu is not None: need_free_mem_indexes = torch.cat((need_free_mem_indexes, eagle_mem_indexes_cpu), dim=0) + self._update_spec_accept_ratio(decode_reqs=decode_reqs, spec_accept_len_cpu=spec_accept_len_cpu) select_mask = torch.tensor(accepted_index_cpu, dtype=torch.bool, device="cpu") self._post_handle( run_reqs=verify_ok_reqs, @@ -896,9 +947,11 @@ def _draft_decode_vanilla_overlap( self, model_input0: ModelInput, model_input1: ModelInput, + model_output0: ModelOutput, + model_output1: ModelOutput, b_req_idx: torch.Tensor, next_token_ids: torch.Tensor = None, - mtp_accept_len: torch.Tensor = None, + spec_accept_len: torch.Tensor = None, b_req_mtp_start_loc: torch.Tensor = None, req_num0: int = 0, req_num1: int = 0, @@ -907,39 +960,43 @@ def _draft_decode_vanilla_overlap( all_next_token_ids.append(next_token_ids) # share some inference info with the main model draft_model_input0, draft_model_input1 = model_input0, model_input1 + draft_hidden0 = model_output0.spec_hidden + draft_hidden1 = model_output1.spec_hidden + assert draft_hidden0 is not None and draft_hidden1 is not None - draft_next_token_ids_gpu0 = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num0 > 0 else None, + draft_next_token_ids_gpu0 = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input0.batch_size, copy_len=req_num0, source_start=0, - device=model_input0.b_req_idx.device, ) - draft_next_token_ids_gpu1 = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num1 > 0 else None, + draft_next_token_ids_gpu1 = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input1.batch_size, copy_len=req_num1, source_start=req_num0, - device=model_input1.b_req_idx.device, ) # process the draft model output - for draft_model_idx in range(self.mtp_step): + for draft_model_idx in range(self.max_draft_step): - draft_model_input0 = self.spec_adapter.prepare_draft_decode_input( + draft_model_input0 = self.spec_engine.prepare_draft_decode_input( model_input=draft_model_input0, next_token_ids=draft_next_token_ids_gpu0, - microbatch_index=0, + mtp_draft_input_hiddens=draft_hidden0, ) - draft_model_input1 = self.spec_adapter.prepare_draft_decode_input( + draft_model_input1 = self.spec_engine.prepare_draft_decode_input( model_input=draft_model_input1, next_token_ids=draft_next_token_ids_gpu1, - microbatch_index=1, + mtp_draft_input_hiddens=draft_hidden1, ) draft_model_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( draft_model_input0, draft_model_input1 ) + draft_hidden0 = draft_model_output0.spec_hidden + draft_hidden1 = draft_model_output1.spec_hidden + assert draft_hidden0 is not None and draft_hidden1 is not None draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) @@ -949,11 +1006,11 @@ def _draft_decode_vanilla_overlap( all_next_token_ids.append(draft_next_token_ids) if req_num0 + req_num1 > 0: - self.spec_adapter.scatter_token_id_steps( + self.spec_engine.scatter_token_id_steps( token_id_steps=all_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, + spec_accept_len=spec_accept_len, ) return None @@ -961,95 +1018,73 @@ def _draft_decode_eagle_overlap( self, model_input0: ModelInput, model_input1: ModelInput, + model_output0: ModelOutput, + model_output1: ModelOutput, b_req_idx: torch.Tensor, next_token_ids: torch.Tensor = None, - mtp_accept_len: torch.Tensor = None, + spec_accept_len: torch.Tensor = None, b_req_mtp_start_loc: torch.Tensor = None, req_num0: int = 0, req_num1: int = 0, ): - all_next_token_ids = [] - all_next_token_ids.append(next_token_ids) - # share some inference info with the main model - draft_model_input0, draft_model_input1 = model_input0, model_input1 - - draft_next_token_ids_gpu0 = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num0 > 0 else None, + verify_width = self.max_draft_step + 1 + assert req_num0 % verify_width == 0 + assert req_num1 % verify_width == 0 + assert model_input0.batch_size % verify_width == 0 + assert model_input1.batch_size % verify_width == 0 + real_request_num0 = req_num0 // verify_width + real_request_num1 = req_num1 // verify_width + request_capacity0 = model_input0.batch_size // verify_width + request_capacity1 = model_input1.batch_size // verify_width + + padded_next_token_ids0 = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input0.batch_size, copy_len=req_num0, source_start=0, - device=model_input0.b_req_idx.device, ) - draft_next_token_ids_gpu1 = self.spec_adapter.build_padded_next_token_ids( - token_ids=next_token_ids if req_num1 > 0 else None, + padded_next_token_ids1 = self._build_padded_next_token_ids( + token_ids=next_token_ids, batch_size=model_input1.batch_size, copy_len=req_num1, source_start=req_num0, + ) + padded_accept_len0 = torch.ones( + (request_capacity0,), + dtype=torch.int32, + device=model_input0.b_req_idx.device, + ) + padded_accept_len1 = torch.ones( + (request_capacity1,), + dtype=torch.int32, device=model_input1.b_req_idx.device, ) - real_req_num0 = req_num0 // (self.mtp_step + 1) - real_req_num1 = req_num1 // (self.mtp_step + 1) - real_req_num = real_req_num0 + real_req_num1 - padded_req_num0 = model_input0.batch_size // (self.mtp_step + 1) - real_req_num0 - padded_req_num1 = model_input1.batch_size // (self.mtp_step + 1) - real_req_num1 - eagle_mem_indexes_cpu = self.spec_adapter.alloc_extra_mem_indexes(real_req_num * self.mtp_step) - eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) - eagle_mem_indexes0 = eagle_mem_indexes[0 : real_req_num0 * self.mtp_step] - eagle_mem_indexes1 = eagle_mem_indexes[real_req_num0 * self.mtp_step : real_req_num * self.mtp_step] - - # process the draft model output - for _step in range(self.mtp_step): - - draft_model_input0 = self.spec_adapter.prepare_draft_decode_input( - model_input=draft_model_input0, - next_token_ids=draft_next_token_ids_gpu0, - microbatch_index=0, - ) - draft_model_input1 = self.spec_adapter.prepare_draft_decode_input( - model_input=draft_model_input1, - next_token_ids=draft_next_token_ids_gpu1, - microbatch_index=1, + if real_request_num0 > 0: + padded_accept_len0[:real_request_num0].copy_(spec_accept_len[:real_request_num0]) + if real_request_num1 > 0: + padded_accept_len1[:real_request_num1].copy_( + spec_accept_len[real_request_num0 : real_request_num0 + real_request_num1] ) - draft_model_idx = _step % self.num_mtp_models - draft_model_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( - draft_model_input0, draft_model_input1 - ) - - draft_model_input0.b_seq_len += 1 - draft_model_input0.max_kv_seq_len += 1 - draft_model_input0.mem_indexes = self.spec_adapter.append_padded_eagle_step_mem_indexes( - model_input=draft_model_input0, - eagle_mem_indexes=eagle_mem_indexes0, - step=_step, - real_req_num=real_req_num0, - padded_req_num=padded_req_num0, - mtp_step=self.mtp_step, - ) - - draft_model_input1.b_seq_len += 1 - draft_model_input1.max_kv_seq_len += 1 - draft_model_input1.mem_indexes = self.spec_adapter.append_padded_eagle_step_mem_indexes( - model_input=draft_model_input1, - eagle_mem_indexes=eagle_mem_indexes1, - step=_step, - real_req_num=real_req_num1, - padded_req_num=padded_req_num1, - mtp_step=self.mtp_step, - ) - - draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) - draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) - draft_next_token_ids = torch.cat( - (draft_next_token_ids_gpu0[0:req_num0], draft_next_token_ids_gpu1[0:req_num1]), dim=0 - ) - all_next_token_ids.append(draft_next_token_ids) + proposal = self.spec_engine.propose_next_overlap( + main_model_input0=model_input0, + main_model_output0=model_output0, + next_token_ids0=padded_next_token_ids0, + real_verify_rows0=req_num0, + accept_len0=padded_accept_len0, + main_model_input1=model_input1, + main_model_output1=model_output1, + next_token_ids1=padded_next_token_ids1, + real_verify_rows1=req_num1, + accept_len1=padded_accept_len1, + draft_step=self.max_draft_step, + ) if req_num0 + req_num1 > 0: - self.spec_adapter.scatter_token_id_steps( - token_id_steps=all_next_token_ids, + self.spec_engine.scatter_next_tokens( b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=proposal.token_ids, b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, + spec_accept_len=spec_accept_len, ) - return eagle_mem_indexes_cpu + return proposal.extra_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py index 68af30b505..d35737e95a 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py @@ -6,9 +6,18 @@ from typing import List, Optional, Tuple from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq from lightllm.utils.infer_utils import calculate_time -from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.envs_utils import ( + enable_diverse_mode_gqa_decode_fast_kernel, + enable_dynamic_spec, + enable_triton_mtp_kernel, + get_env_start_args, +) from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from .generic_pre_process import build_b_position_delta +from .generic_pre_process import ( + build_b_position_delta, + build_diverse_shared_group_infos, + build_spec_shared_group_markers, +) def padded_prepare_prefill_inputs( @@ -150,7 +159,7 @@ def padded_prepare_decode_inputs( b_mtp_index = [] b_seq_len = [] b_q_seq_len = [] - args_mtp_step = get_env_start_args().mtp_step + draft_step = get_env_start_args().mtp_step batch_multimodal_params = [] for req in req_objs: run_reqs.append(req) @@ -163,7 +172,7 @@ def padded_prepare_decode_inputs( b_mtp_index.append(0) batch_multimodal_params.append(req.multimodal_params) # process the draft tokens. - for step in range(req.mtp_step): + for step in range(req.max_draft_step): run_reqs.append(req) seq_len += 1 total_token_num += seq_len @@ -182,7 +191,7 @@ def padded_prepare_decode_inputs( b_q_seq_len.append(1) b_mtp_index.append(0) batch_multimodal_params.append({"images": [], "audios": []}) - for step in range(args_mtp_step): + for step in range(draft_step): seq_len += 1 total_token_num += seq_len b_seq_len.append(seq_len) @@ -199,16 +208,28 @@ def padded_prepare_decode_inputs( b_mtp_index = torch.tensor(b_mtp_index, dtype=torch.int32, device="cpu") b_position_delta = build_b_position_delta(batch_multimodal_params) + padded_row_count = padded_req_num * (draft_step + 1) + if enable_diverse_mode_gqa_decode_fast_kernel(): + b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) + if padded_row_count > 0: + b_shared_seq_len = F.pad(b_shared_seq_len, (0, padded_row_count), value=0) + b_mark_shared_group = F.pad(b_mark_shared_group, (0, padded_row_count), value=1) + elif enable_dynamic_spec() or enable_triton_mtp_kernel(): + b_shared_seq_len = None + b_mark_shared_group = build_spec_shared_group_markers(b_mtp_index=b_mtp_index) + else: + b_shared_seq_len = None + b_mark_shared_group = None + # dynamic prompt cache 准备 token - padded_mem_indexes_num = padded_req_num * (args_mtp_step + 1) if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(b_seq_len.shape[0] - padded_mem_indexes_num) - mem_indexes = g_infer_context.req_manager.mem_manager.alloc(b_seq_len.shape[0] - padded_mem_indexes_num) + g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(b_seq_len.shape[0] - padded_row_count) + mem_indexes = g_infer_context.req_manager.mem_manager.alloc(b_seq_len.shape[0] - padded_row_count) - if padded_mem_indexes_num > 0: + if padded_row_count > 0: mem_indexes = F.pad( input=mem_indexes, - pad=(0, padded_mem_indexes_num), + pad=(0, padded_row_count), mode="constant", value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, ) @@ -224,6 +245,9 @@ def padded_prepare_decode_inputs( b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, b_position_delta=b_position_delta, + b_shared_seq_len=b_shared_seq_len, + b_mark_shared_group=b_mark_shared_group, + draft_step=draft_step, is_prefill=False, multimodal_params=batch_multimodal_params, ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py index 883f444e41..d0e035aed0 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py @@ -1,28 +1,20 @@ -from __future__ import annotations - import torch -import triton -import triton.language as tl -from typing import TYPE_CHECKING, List, Tuple, Optional +from typing import List, Tuple, Optional +from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import trim_post_sample_tensors from lightllm.common.basemodel.triton_kernel.post_process.apply_penalty import apply_penalty from lightllm.common.basemodel.triton_kernel.post_process.apply_penalty_gpu_cache import apply_penalty_gpu_cache from lightllm.common.basemodel.triton_kernel.post_process.apply_invalid_token import apply_invalid_token_ids +from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.utils.envs_utils import get_env_start_args -if TYPE_CHECKING: - from lightllm.server.router.model_infer.infer_batch import InferReq - def sample( logits: torch.Tensor, reqs: List[InferReq], eos_id: List[int] = [2], - dynamic_batch_size: Optional[int] = None, - selected_run_reqs: Optional[torch.Tensor] = None, + selected_row_mask: Optional[torch.Tensor] = None, ): - from lightllm.server.router.model_infer.infer_batch import g_infer_context - from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager - ( b_req_idx, b_temperatures, @@ -39,9 +31,9 @@ def sample( exist_req_use_random_seed, ) = _get_post_sample_tensors(reqs) + sampling_params_manager = g_infer_context.req_manager.req_sampling_params_manager sample_reqs = reqs - if selected_run_reqs is not None: - assert dynamic_batch_size is not None + if selected_row_mask is not None: ( b_req_idx, b_temperatures, @@ -49,9 +41,9 @@ def sample( b_top_ks, b_length_penalty_param, b_mask_eos_reqs, - ) = _trim_post_sample_tensors( - dynamic_batch_size=dynamic_batch_size, - selected_run_reqs=selected_run_reqs, + ) = trim_post_sample_tensors( + dynamic_batch_size=logits.shape[0], + selected_row_mask=selected_row_mask, b_req_idx=b_req_idx, b_temperatures=b_temperatures, b_top_ps=b_top_ps, @@ -59,8 +51,12 @@ def sample( b_length_penalty_param=b_length_penalty_param, b_mask_eos_reqs=b_mask_eos_reqs, ) - if has_invalid_token_ids or exist_req_use_random_seed: - sample_reqs = _get_selected_reqs(reqs=reqs, selected_run_reqs=selected_run_reqs) + if ( + has_invalid_token_ids + or exist_req_use_random_seed + or sampling_params_manager.penalty_counter_mode == "cpu_counter" + ): + sample_reqs = _get_selected_reqs(reqs=reqs, selected_row_mask=selected_row_mask) if has_invalid_token_ids: invalid_token_ids, cu_invalid_token_num, has_invalid_token_ids = _get_invalid_token_tensors( reqs=sample_reqs @@ -70,8 +66,6 @@ def sample( eos_ids = g_pin_mem_manager.gen_from_list(key="eos_ids", data=eos_id, dtype=torch.int32).cuda(non_blocking=True) - sampling_params_manager = g_infer_context.req_manager.req_sampling_params_manager - # 这里需要区分历史token的频率惩罚类的系数的生效模式,目前支持两种在线统计方式: # 一种是基于 cpu 的,每个 req 对象利用其上绑定的dict对象out_token_id_count,每生成一个token就进行相应 # 的计数更新,当进行使用的时候, 对一个需要处理的req list, 会生成对应的3个 triton kernel 需要使用的惩罚系数 @@ -89,7 +83,7 @@ def sample( p_token_ids, p_token_counts, p_cumsum_seq_len, - ) = sampling_params_manager.gen_cpu_out_token_counter_sampling_params(req_objs=reqs) + ) = sampling_params_manager.gen_cpu_out_token_counter_sampling_params(req_objs=sample_reqs) apply_penalty( Logits=logits, @@ -199,8 +193,6 @@ def _random_sample(probs: torch.Tensor, reqs: List[InferReq], exist_req_use_rand def _get_post_sample_tensors(reqs: List[InferReq]): - from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager - req_idxes: List[int] = [] temperatures: List[float] = [] top_ps: List[float] = [] @@ -279,14 +271,12 @@ def _get_post_sample_tensors(reqs: List[InferReq]): ) -def _get_selected_reqs(reqs: List[InferReq], selected_run_reqs: torch.Tensor): - selected_run_reqs_cpu = selected_run_reqs.detach().cpu().tolist() - return [req_obj for req_obj, selected in zip(reqs, selected_run_reqs_cpu) if int(selected) != 0] +def _get_selected_reqs(reqs: List[InferReq], selected_row_mask: torch.Tensor): + selected_row_mask_cpu = selected_row_mask.detach().cpu().tolist() + return [req_obj for req_obj, selected in zip(reqs, selected_row_mask_cpu) if int(selected) != 0] def _get_invalid_token_tensors(reqs: List[InferReq]): - from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager - invalid_token_ids: List[int] = [] has_invalid_token_ids = False cu_invalid_token_num = [0] @@ -313,102 +303,3 @@ def _get_invalid_token_tensors(reqs: List[InferReq]): cu_invalid_token_num_cpu.cuda(non_blocking=True), True, ) - - -@triton.jit -def _fwd_kernel_trim_post_sample_tensors( - b_req_idx, - out_b_req_idx, - b_temperatures, - out_b_temperatures, - b_top_ps, - out_b_top_ps, - b_top_ks, - out_b_top_ks, - b_length_penalty_param, - out_b_length_penalty_param, - b_mask_eos_reqs, - out_b_mask_eos_reqs, - selected_run_reqs, - selected_dst_pos, - batch_size, - BLOCK_SIZE: tl.constexpr, -): - pid = tl.program_id(0) - offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offsets < batch_size - selected = tl.load(selected_run_reqs + offsets, mask=mask, other=0) != 0 - dst_pos = tl.load(selected_dst_pos + offsets, mask=mask, other=0) - write_mask = mask & selected - - req_idx = tl.load(b_req_idx + offsets, mask=mask, other=0) - temperature = tl.load(b_temperatures + offsets, mask=mask, other=0.0) - top_p = tl.load(b_top_ps + offsets, mask=mask, other=0.0) - top_k = tl.load(b_top_ks + offsets, mask=mask, other=0) - length_penalty = tl.load(b_length_penalty_param + offsets, mask=mask, other=0) - mask_eos_req = tl.load(b_mask_eos_reqs + offsets, mask=mask, other=0) - - tl.store(out_b_req_idx + dst_pos, req_idx, mask=write_mask) - tl.store(out_b_temperatures + dst_pos, temperature, mask=write_mask) - tl.store(out_b_top_ps + dst_pos, top_p, mask=write_mask) - tl.store(out_b_top_ks + dst_pos, top_k, mask=write_mask) - tl.store(out_b_length_penalty_param + dst_pos, length_penalty, mask=write_mask) - tl.store(out_b_mask_eos_reqs + dst_pos, mask_eos_req, mask=write_mask) - - -def _trim_post_sample_tensors( - dynamic_batch_size: int, - selected_run_reqs: torch.Tensor, - b_req_idx: torch.Tensor, - b_temperatures: torch.Tensor, - b_top_ps: torch.Tensor, - b_top_ks: torch.Tensor, - b_length_penalty_param: torch.Tensor, - b_mask_eos_reqs: torch.Tensor, -): - assert selected_run_reqs.is_cuda - dynamic_batch_size = int(dynamic_batch_size) - selected_run_reqs = selected_run_reqs.to(torch.int32) - selected_dst_pos = torch.cumsum(selected_run_reqs, dim=0, dtype=torch.int32) - 1 - batch_size = selected_run_reqs.shape[0] - - out_b_req_idx = torch.empty((dynamic_batch_size,), dtype=b_req_idx.dtype, device=b_req_idx.device) - out_b_temperatures = torch.empty((dynamic_batch_size,), dtype=b_temperatures.dtype, device=b_temperatures.device) - out_b_top_ps = torch.empty((dynamic_batch_size,), dtype=b_top_ps.dtype, device=b_top_ps.device) - out_b_top_ks = torch.empty((dynamic_batch_size,), dtype=b_top_ks.dtype, device=b_top_ks.device) - out_b_length_penalty_param = torch.empty( - (dynamic_batch_size,), dtype=b_length_penalty_param.dtype, device=b_length_penalty_param.device - ) - out_b_mask_eos_reqs = torch.empty((dynamic_batch_size,), dtype=b_mask_eos_reqs.dtype, device=b_mask_eos_reqs.device) - - BLOCK_SIZE = 256 - grid = (triton.cdiv(batch_size, BLOCK_SIZE),) - _fwd_kernel_trim_post_sample_tensors[grid]( - b_req_idx=b_req_idx, - out_b_req_idx=out_b_req_idx, - b_temperatures=b_temperatures, - out_b_temperatures=out_b_temperatures, - b_top_ps=b_top_ps, - out_b_top_ps=out_b_top_ps, - b_top_ks=b_top_ks, - out_b_top_ks=out_b_top_ks, - b_length_penalty_param=b_length_penalty_param, - out_b_length_penalty_param=out_b_length_penalty_param, - b_mask_eos_reqs=b_mask_eos_reqs, - out_b_mask_eos_reqs=out_b_mask_eos_reqs, - selected_run_reqs=selected_run_reqs, - selected_dst_pos=selected_dst_pos, - batch_size=batch_size, - BLOCK_SIZE=BLOCK_SIZE, - num_warps=4, - num_stages=1, - ) - - return ( - out_b_req_idx, - out_b_temperatures, - out_b_top_ps, - out_b_top_ks, - out_b_length_penalty_param, - out_b_mask_eos_reqs, - ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index db7414e85a..e031b5e318 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -7,7 +7,8 @@ enable_diverse_mode_gqa_decode_fast_kernel, enable_triton_mtp_kernel, get_diverse_max_batch_shared_group_size, - enable_dynamic_mtp_verify, + enable_dynamic_spec, + get_env_start_args, ) @@ -116,7 +117,7 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_mtp_index.append(0) multimodal_params.append(req.multimodal_params) # process the draft tokens. - for step in range(req.mtp_step): + for step in range(req.max_draft_step): run_reqs.append(req) b_req_idx.append(req.req_idx) seq_len += 1 @@ -134,13 +135,11 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_mtp_index = torch.tensor(b_mtp_index, dtype=torch.int32, device="cpu") b_position_delta = build_b_position_delta(multimodal_params) - # diverse mode 和 dynamic MTP mode 使用不同的 shared group 构建逻辑 if enable_diverse_mode_gqa_decode_fast_kernel(): b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) - elif enable_dynamic_mtp_verify() or enable_triton_mtp_kernel(): - # MTP 模式下,使用专门的 shared group 构建函数 - b_shared_seq_len = None # MTP 模式不需要 b_shared_seq_len - b_mark_shared_group = build_mtp_shared_group_infos(run_reqs=run_reqs) + elif enable_dynamic_spec() or enable_triton_mtp_kernel(): + b_shared_seq_len = None + b_mark_shared_group = build_spec_shared_group_markers(b_mtp_index=b_mtp_index) else: b_shared_seq_len = None b_mark_shared_group = None @@ -163,6 +162,7 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_position_delta=b_position_delta, b_shared_seq_len=b_shared_seq_len, b_mark_shared_group=b_mark_shared_group, + draft_step=get_env_start_args().mtp_step, is_prefill=False, multimodal_params=multimodal_params, ) @@ -226,33 +226,23 @@ def build_diverse_shared_group_infos(run_reqs: List[InferReq]) -> Tuple[torch.Te return b_shared_seq_len, b_mark_shared_group -def build_mtp_shared_group_infos(run_reqs: List[InferReq]) -> torch.Tensor: - # Similar to build_diverse_shared_group_infos, - # but the grouping logic is based on b_mtp_index, which indicates the MTP step of each request +def build_spec_shared_group_markers(b_mtp_index: torch.Tensor) -> torch.Tensor: + # Each logical request starts at row index 0. Only the final row of each + # speculative query group stores its group size; earlier rows store zero. max_batch_shared_group_size = get_diverse_max_batch_shared_group_size() - req_ids = [req.req_id for req in run_reqs] b_mark_shared_group = [] - _current_group = [] - for node in req_ids: - if not _current_group: - _current_group.append(node) - elif node == _current_group[-1]: - _current_group.append(node) - else: - b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) - b_mark_shared_group[-1] = len(_current_group) - _current_group.clear() - _current_group.append(node) - - if len(_current_group) == max_batch_shared_group_size: - b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) - b_mark_shared_group[-1] = len(_current_group) - _current_group.clear() - if _current_group: - b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) - b_mark_shared_group[-1] = len(_current_group) - _current_group.clear() - - assert len(b_mark_shared_group) == len(run_reqs) + group_start = 0 + for group_end in range(1, len(b_mtp_index) + 1): + reaches_request_boundary = group_end == len(b_mtp_index) or b_mtp_index[group_end] == 0 + reaches_size_limit = group_end - group_start == max_batch_shared_group_size + if not reaches_request_boundary and not reaches_size_limit: + continue + + group_size = group_end - group_start + b_mark_shared_group.extend([0] * (group_size - 1)) + b_mark_shared_group.append(group_size) + group_start = group_end + + assert len(b_mark_shared_group) == len(b_mtp_index) b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cpu") return b_mark_shared_group diff --git a/lightllm/server/router/model_infer/mode_backend/update_mem_index.py b/lightllm/server/router/model_infer/mode_backend/update_mem_index.py deleted file mode 100644 index 4b0c7c441e..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/update_mem_index.py +++ /dev/null @@ -1,50 +0,0 @@ -import torch -import triton -import triton.language as tl - - -@triton.jit -def update_eagle_mem_indexes_kernel( - old_mems_ptr, # [N] - new_step_mems_ptr, # [num_reqs] - b_req_mtp_start_loc, # [num_reqs] - out_mems_ptr, # [N] - req_all_num, - BLOCK_SIZE: tl.constexpr, -): - cur_req_idx = tl.program_id(0) - origin_req_num = tl.num_programs(0) - - offs = tl.arange(0, BLOCK_SIZE) - start_loc = tl.load(b_req_mtp_start_loc + cur_req_idx) - end_loc = tl.load(b_req_mtp_start_loc + cur_req_idx + 1, mask=cur_req_idx + 1 < origin_req_num, other=req_all_num) - - req_mtp_num = end_loc - start_loc - old_mems = tl.load(old_mems_ptr + start_loc + offs + 1, mask=offs + 1 < req_mtp_num, other=0) - tl.store(out_mems_ptr + start_loc + offs, old_mems, mask=offs + 1 < req_mtp_num) - new_step_mems = tl.load(new_step_mems_ptr + cur_req_idx) - tl.store(out_mems_ptr + end_loc - 1, new_step_mems) - - -def update_eagle_mem_indexes_triton( - old_mem_indexes: torch.Tensor, new_step_mem_indexes: torch.Tensor, b_req_mtp_start_loc: torch.Tensor -): - """ - old_mem_indexes: [N] CUDA Tensor - new_step_mem_indexes: [num_reqs] CUDA Tensor - """ - out = torch.empty_like(old_mem_indexes) - BLOCK_SIZE = 32 - original_num_reqs = b_req_mtp_start_loc.shape[0] - assert original_num_reqs == new_step_mem_indexes.shape[0] - req_all_num = old_mem_indexes.shape[0] - grid = (original_num_reqs,) - update_eagle_mem_indexes_kernel[grid]( - old_mems_ptr=old_mem_indexes, - new_step_mems_ptr=new_step_mem_indexes, - b_req_mtp_start_loc=b_req_mtp_start_loc, - out_mems_ptr=out, - req_all_num=req_all_num, - BLOCK_SIZE=BLOCK_SIZE, - ) - return out diff --git a/lightllm/server/router/model_infer/speculative/__init__.py b/lightllm/server/router/model_infer/speculative/__init__.py index 6a238ecf00..0697431885 100644 --- a/lightllm/server/router/model_infer/speculative/__init__.py +++ b/lightllm/server/router/model_infer/speculative/__init__.py @@ -1,42 +1,4 @@ -__all__ = [ - "SpecRuntime", - "SpecDecodeForwardState", - "SpecDecodePostState", - "SpecDecodeRunner", - "SpecVerifier", - "SpecVerifyResult", - "build_spec_runtime", -] +from lightllm.server.router.model_infer.speculative.engine import SpecEngine -def __getattr__(name): - if name in ("SpecRuntime", "build_spec_runtime"): - from lightllm.server.router.model_infer.speculative.runtime import SpecRuntime, build_spec_runtime - - values = { - "SpecRuntime": SpecRuntime, - "build_spec_runtime": build_spec_runtime, - } - return values[name] - if name in ("SpecDecodeForwardState", "SpecDecodePostState", "SpecDecodeRunner"): - from lightllm.server.router.model_infer.speculative.runner import ( - SpecDecodeForwardState, - SpecDecodePostState, - SpecDecodeRunner, - ) - - values = { - "SpecDecodeForwardState": SpecDecodeForwardState, - "SpecDecodePostState": SpecDecodePostState, - "SpecDecodeRunner": SpecDecodeRunner, - } - return values[name] - if name in ("SpecVerifier", "SpecVerifyResult"): - from lightllm.server.router.model_infer.speculative.verifier import SpecVerifier, SpecVerifyResult - - values = { - "SpecVerifier": SpecVerifier, - "SpecVerifyResult": SpecVerifyResult, - } - return values[name] - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") +__all__ = ["SpecEngine"] diff --git a/lightllm/server/router/model_infer/speculative/engine.py b/lightllm/server/router/model_infer/speculative/engine.py new file mode 100644 index 0000000000..d4130641a5 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/engine.py @@ -0,0 +1,531 @@ +from __future__ import annotations + +from typing import Callable, List, Optional, Tuple + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +from lightllm.server.router.model_infer.speculative.planner import ( + DSparkDynamicSpecPlanner, + DynamicSpecPlanner, + Eagle3DynamicSpecPlanner, + FixedSpecPlanner, + SpecDecodePlan, +) +from lightllm.server.router.model_infer.speculative.proposers import build_spec_proposer +from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.speculative.runner import ( + SpecDecodeForwardState, + SpecDecodePostState, + SpecDecodeRunner, +) +from lightllm.server.router.model_infer.speculative.verifier import SpecVerifier, SpecVerifyResult +from lightllm.utils.envs_utils import enable_dynamic_spec + + +class SpecEngine: + """Owns one speculative decoding pipeline for a model backend. + + BaseModel only returns ``ModelOutput.spec_hidden``. The engine owns + planning, target verification, draft extension/decoding, and the request + bookkeeping around that pipeline. Algorithm-specific state stays in the + proposer. + + The target-to-draft data path is explicit: + 1. target forward returns ``ModelOutput.spec_hidden`` + 2. proposer performs one draft extend from that hidden state + 3. recurrent proposers perform zero or more unit-stride draft decodes + 4. verifier accepts target rows and scatters the next candidate block + """ + + def __init__(self, backend) -> None: + self.backend = backend + self.spec_mode = backend.args.mtp_mode + self.enable_dynamic_spec = enable_dynamic_spec() + self.verifier = SpecVerifier(backend=backend) + self.proposer = build_spec_proposer(engine=self) + self.decode_runner = SpecDecodeRunner(engine=self) + self.planner = self._build_decode_planner() + self._register_cuda_graph_costs() + self._dynamic_accept_stats_calls = 0 + + def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: + """Allocate speculative draft-owned temporary KV slots.""" + + token_count = int(token_count) + assert token_count >= 0 + if token_count == 0: + return torch.empty((0,), dtype=torch.int32, device="cpu") + + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + if g_infer_context.radix_cache is not None: + g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) + return g_infer_context.req_manager.mem_manager.alloc(token_count) + + def prepare_draft_prefill_input( + self, + model_input: ModelInput, + next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, + ) -> ModelInput: + """Build draft prefill input from target prefill input. + + `next_token_ids`: [run_req_num] + `mtp_draft_input_hiddens`: captured target feature. The first + dimension matches the target prefill token layout after padding/unpad + handling; the second dimension is either hidden_size or + hidden_size * len(target_layer_ids). + """ + + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + return prepare_mtp_prefill_inputs( + model_input=model_input, + b_next_token_ids=next_token_ids, + mtp_draft_input_hiddens=mtp_draft_input_hiddens, + ) + + def prepare_draft_decode_input( + self, + model_input: ModelInput, + next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, + ) -> ModelInput: + """Mutate a decode ModelInput for one draft forward. + + `next_token_ids`: [verify_batch] + `mtp_draft_input_hiddens`: [verify_batch, hidden_dim_for_draft] + """ + + model_input.input_ids = next_token_ids + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + if self.spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3", "qwen3next_eagle"): + model_input.draft_step = 0 + return model_input + + def build_initial_draft_state( + self, + model_input: ModelInput, + model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + self.proposer.build_initial_draft_state( + model_input=model_input, + model_output=model_output, + next_token_ids=next_token_ids, + ) + + def build_initial_draft_state_overlap( + self, + model_input0: ModelInput, + model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + model_input1: ModelInput, + model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + ) -> None: + self.proposer.build_initial_draft_state_overlap( + model_input0=model_input0, + model_output0=model_output0, + next_token_ids0=next_token_ids0, + model_input1=model_input1, + model_output1=model_output1, + next_token_ids1=next_token_ids1, + ) + + def plan_decode(self, model_input: ModelInput, req_num: int) -> SpecDecodePlan: + """Return the fixed or dynamic speculative plan for one decode iteration.""" + + return self.planner.plan(req_num=req_num, original_batch_size=model_input.batch_size) + + def run_decode_speculative_forward( + self, + model_input: ModelInput, + model_output: ModelOutput, + run_reqs: List, + req_num: int, + plan: SpecDecodePlan, + selected_row_mask_cpu: Optional[torch.Tensor], + next_token_ids: torch.Tensor, + next_token_logprobs: torch.Tensor, + next_token_ranks: torch.Tensor, + copy_next_token_infos: Callable[ + [torch.Tensor, torch.Tensor, torch.Tensor], + Tuple[torch.Tensor, torch.Tensor, torch.Tensor], + ], + ) -> SpecDecodeForwardState: + return self.decode_runner.run_speculative_forward( + model_input=model_input, + model_output=model_output, + run_reqs=run_reqs, + req_num=req_num, + plan=plan, + selected_row_mask_cpu=selected_row_mask_cpu, + next_token_ids=next_token_ids, + next_token_logprobs=next_token_logprobs, + next_token_ranks=next_token_ranks, + copy_next_token_infos=copy_next_token_infos, + ) + + def resolve_decode_pre_post_reqs(self, state: SpecDecodeForwardState, decode_reqs: List): + return self.decode_runner.resolve_pre_post_reqs(state=state, decode_reqs=decode_reqs) + + def finish_decode_post( + self, + state: SpecDecodeForwardState, + req_num: int, + run_reqs: List, + ) -> SpecDecodePostState: + return self.decode_runner.finish_post(state=state, req_num=req_num, run_reqs=run_reqs) + + def prepare_decode_model_input( + self, + model_input: ModelInput, + req_num: int, + plan: SpecDecodePlan, + ): + """Apply target verify-row compaction when the dynamic planner selects it.""" + + if not plan.is_dynamic: + return model_input, None + + if plan.dynamic_batch_size == model_input.batch_size: + return model_input, None + + self._clear_stale_dynamic_token_probs(pre_draft_step=plan.pre_draft_step) + + from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_spec_model_input + + model_input, selected_row_mask = prepare_dynamic_spec_model_input( + model_input=model_input, + req_num=req_num, + dynamic_batch_size=plan.dynamic_batch_size, + req_to_next_token_probs=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs, + pre_draft_step=plan.pre_draft_step, + ) + return model_input, selected_row_mask + + def async_copy_selected_row_mask(self, selected_row_mask: Optional[torch.Tensor]): + if selected_row_mask is None: + return None + return g_pin_mem_manager.async_copy_from_gpu_tensor( + key="selected_row_mask", + gpu_tensor=selected_row_mask, + ) + + def build_decode_req_lists( + self, + original_run_reqs, + selected_row_mask_cpu: Optional[torch.Tensor], + accepted_index_cpu: torch.Tensor, + ): + """Build post-handle request lists after optional verify-row compaction.""" + + if self.enable_dynamic_spec and selected_row_mask_cpu is not None: + selected_row_mask_numpy = selected_row_mask_cpu.numpy() + run_reqs = [original_run_reqs[i] for i in range(len(original_run_reqs)) if selected_row_mask_numpy[i] == 1] + else: + run_reqs = original_run_reqs + + accepted_index_cpu_numpy = accepted_index_cpu.numpy() + verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu_numpy[i] == 1] + return run_reqs, verify_ok_reqs + + def build_decode_free_mem_indexes_cpu( + self, + model_input: ModelInput, + selected_row_mask_cpu: Optional[torch.Tensor], + accepted_index_cpu: torch.Tensor, + ) -> torch.Tensor: + mem_indexes_cpu = model_input.mem_indexes_cpu + if not self.enable_dynamic_spec or selected_row_mask_cpu is None: + return mem_indexes_cpu[accepted_index_cpu == 0] + + selected_mask = selected_row_mask_cpu.to(dtype=torch.bool) + accepted_mask = accepted_index_cpu.to(dtype=torch.bool) + selected_mem_indexes_cpu = mem_indexes_cpu[selected_mask] + assert selected_mem_indexes_cpu.shape[0] == accepted_mask.shape[0] + + unselected_mem_indexes_cpu = mem_indexes_cpu[~selected_mask] + rejected_selected_mem_indexes_cpu = selected_mem_indexes_cpu[~accepted_mask] + if len(unselected_mem_indexes_cpu) == 0: + return rejected_selected_mem_indexes_cpu + if len(rejected_selected_mem_indexes_cpu) == 0: + return unselected_mem_indexes_cpu + return torch.cat([unselected_mem_indexes_cpu, rejected_selected_mem_indexes_cpu], dim=0) + + def update_dynamic_accept_stats( + self, + req_num: int, + run_reqs, + accepted_index_cpu: torch.Tensor, + spec_accept_len_cpu: torch.Tensor, + dynamic_batch_size: Optional[int], + pre_draft_step: Optional[int] = None, + ) -> None: + if not self.enable_dynamic_spec: + return + + assert dynamic_batch_size is not None + assert len(run_reqs) == accepted_index_cpu.shape[0] + assert spec_accept_len_cpu.shape[0] == req_num + accept_lengths = spec_accept_len_cpu.numpy() + accept_count = int(accept_lengths.sum()) + total_count = int(dynamic_batch_size) + is_full_verify = dynamic_batch_size == req_num * (self.backend.max_draft_step + 1) + + # The first target decode has no preceding draft proposal, so every + # request structurally accepts only its base token. Treating that + # cold-start iteration as a K-wide acceptance sample makes a highly + # predictable workload look maximally hard and can collapse the + # controller before the first real proposal is verified. + self._dynamic_accept_stats_calls += 1 + if self._dynamic_accept_stats_calls == 1: + return + + if self.spec_mode == "eagle3": + self.planner.update_req_num_to_dynamic_batch_size_to_accept_ratio( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + accept_ratio=accept_count / total_count, + pre_draft_step=pre_draft_step, + ) + if is_full_verify: + self.planner.update_full_verify_tokens_per_req( + accept_count / req_num, + req_num=req_num, + ) + self.planner.update_observed_iteration_stats( + tokens_per_req=accept_count / req_num, + verify_rows_per_req=dynamic_batch_size / req_num, + is_full_verify=is_full_verify, + req_num=req_num, + ) + # Confidence-selected rows bias the survival curve upward, so only + # full-width verification provides valid Eagle3 prefix samples. + if is_full_verify: + self.planner.update_verified_batch_prefix_stats( + verify_and_accept_lengths=[ + (self.backend.max_draft_step + 1, int(accept_len)) for accept_len in accept_lengths + ], + ) + else: + self.planner.update_req_num_to_dynamic_batch_size_to_accept_ratio( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + accept_ratio=accept_count / total_count, + ) + verify_rows_per_req = max(1, int(round(dynamic_batch_size / req_num))) + for accept_len in accept_lengths: + self.planner.update_verified_prefix_stats( + verify_len=verify_rows_per_req, + accept_len=int(accept_len), + ) + + def needs_schedule_probs_cpu(self) -> bool: + """Whether this planner consumes proposal confidence on the CPU.""" + + return self.enable_dynamic_spec and self.spec_mode == "dspark" + + def update_dynamic_schedule_stats( + self, + req_num: int, + schedule_probs_cpu: Optional[torch.Tensor], + ) -> None: + if not self.needs_schedule_probs_cpu() or schedule_probs_cpu is None: + return + + self.planner.update_predicted_schedule_probs( + schedule_probs=schedule_probs_cpu, + req_num=req_num, + ) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: Optional[ModelOutput], + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: Optional[torch.Tensor] = None, + ) -> SpecProposal: + return self.proposer.propose_next( + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=draft_step, + accept_len=accept_len, + ) + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor, + draft_step: int, + ) -> SpecProposal: + return self.proposer.propose_next_overlap( + main_model_input0=main_model_input0, + main_model_output0=main_model_output0, + next_token_ids0=next_token_ids0, + real_verify_rows0=real_verify_rows0, + accept_len0=accept_len0, + main_model_input1=main_model_input1, + main_model_output1=main_model_output1, + next_token_ids1=next_token_ids1, + real_verify_rows1=real_verify_rows1, + accept_len1=accept_len1, + draft_step=draft_step, + ) + + def verify_target_tokens( + self, + new_next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + ) -> SpecVerifyResult: + return self.verifier.verify_target_tokens( + new_next_token_ids=new_next_token_ids, + b_req_idx=b_req_idx, + b_req_mtp_start_loc=b_req_mtp_start_loc, + ) + + def build_all_next_token_probs( + self, + proposal: SpecProposal, + draft_step: int, + ) -> Optional[torch.Tensor]: + """Build selected-token probabilities for dynamic speculative scatter. + + Output shape is [verify_batch, max_draft_step + 1]. Column 0 is the target + token probability, fixed to 1 because the target sample is always the + base accepted position. Draft columns store selected-token + probabilities from each proposer step. + """ + + if not self.enable_dynamic_spec: + return None + + schedule_probs = proposal.schedule_probs if proposal.schedule_probs is not None else proposal.draft_probs + assert schedule_probs is not None + + verify_row_count = proposal.token_ids.shape[0] + all_next_token_probs = torch.zeros( + size=(verify_row_count, self.backend.max_draft_step + 1), + dtype=torch.float32, + device=proposal.token_ids.device, + ) + all_next_token_probs[:, 0] = 1.0 + + if isinstance(schedule_probs, torch.Tensor): + assert schedule_probs.shape == (verify_row_count, draft_step) + if draft_step > 0: + all_next_token_probs[:, 1 : draft_step + 1] = schedule_probs + return all_next_token_probs + + assert len(schedule_probs) == draft_step + for step_idx, step_probs in enumerate(schedule_probs): + all_next_token_probs[:, step_idx + 1] = step_probs + return all_next_token_probs + + def pad_all_next_token_ids(self, token_ids: torch.Tensor, draft_step: int) -> torch.Tensor: + """Pad a shorter dynamic proposal to the configured speculative width.""" + + if not self.enable_dynamic_spec or draft_step >= self.backend.max_draft_step: + return token_ids + + append_next_token_ids = torch.ones( + size=(token_ids.shape[0], self.backend.max_draft_step - draft_step), + dtype=token_ids.dtype, + device=token_ids.device, + ) + return torch.cat([token_ids, append_next_token_ids], dim=-1) + + def scatter_next_tokens( + self, + b_req_mtp_start_loc: torch.Tensor, + all_next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + spec_accept_len: torch.Tensor, + all_next_token_probs: Optional[torch.Tensor] = None, + ) -> None: + self.verifier.scatter_next_tokens( + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=b_req_idx, + spec_accept_len=spec_accept_len, + all_next_token_probs=all_next_token_probs, + ) + + def scatter_token_id_steps( + self, + token_id_steps: List[torch.Tensor], + b_req_mtp_start_loc: torch.Tensor, + b_req_idx: torch.Tensor, + spec_accept_len: torch.Tensor, + row_count: Optional[int] = None, + ) -> torch.Tensor: + """Stack proposal token columns and scatter them for the next verify.""" + + all_next_token_ids = torch.stack(token_id_steps, dim=1) + if row_count is not None: + all_next_token_ids = all_next_token_ids[: int(row_count), :] + self.scatter_next_tokens( + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=b_req_idx, + spec_accept_len=spec_accept_len, + ) + return all_next_token_ids + + def _build_decode_planner(self): + if not self.enable_dynamic_spec: + return FixedSpecPlanner(max_draft_step=self.backend.max_draft_step) + if self.spec_mode == "dspark": + return DSparkDynamicSpecPlanner(max_draft_step=self.backend.max_draft_step) + if self.spec_mode == "eagle3": + return Eagle3DynamicSpecPlanner(max_draft_step=self.backend.max_draft_step) + return DynamicSpecPlanner(max_draft_step=self.backend.max_draft_step) + + def _register_cuda_graph_costs(self) -> None: + if not self.enable_dynamic_spec: + return + + target_graph = self.backend.model.graph + if target_graph is not None: + for batch_size, infer_cost_ms in target_graph.infer_cost_ms_by_batch_size.items(): + self.planner.update_infer_cost( + batch_size=batch_size, + infer_cost_ms=infer_cost_ms, + is_draft_model=False, + ) + + if self.spec_mode == "dspark": + return + for draft_model in self.backend.draft_models: + draft_graph = draft_model.graph + if draft_graph is None: + continue + for batch_size, infer_cost_ms in draft_graph.infer_cost_ms_by_batch_size.items(): + self.planner.update_infer_cost( + batch_size=batch_size, + infer_cost_ms=infer_cost_ms, + is_draft_model=True, + ) + + def _clear_stale_dynamic_token_probs(self, pre_draft_step: int) -> None: + # Columns after the previous draft length are stale and must not be + # sampled by dynamic row compaction in the current target forward. + self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs[ + :, (pre_draft_step + 1) : + ].fill_(0.0) diff --git a/lightllm/server/router/model_infer/speculative/planner.py b/lightllm/server/router/model_infer/speculative/planner.py index 66969663e2..5781831e96 100644 --- a/lightllm/server/router/model_infer/speculative/planner.py +++ b/lightllm/server/router/model_infer/speculative/planner.py @@ -3,28 +3,23 @@ import math import os import random -from collections import Counter, deque +from collections import deque from dataclasses import dataclass from typing import Dict, List, Optional, Tuple import numpy as np from sortedcontainers import SortedDict -from lightllm.utils.log_utils import init_logger - - -logger = init_logger(__name__) - @dataclass(frozen=True) class SpecDecodePlan: """Planner decision for one target decode iteration. - Static MTP uses the full MTP-expanded target batch: + Fixed scheduling uses the full speculative-expanded target batch: - dynamic_batch_size is None - - draft_step == mtp_step + - draft_step == max_draft_step - Dynamic MTP may compact target rows before forward: + Dynamic speculative scheduling may compact target rows before forward: - dynamic_batch_size is the selected target row count - draft_step is the candidate length to generate after target verify - pre_draft_step describes the previous iteration and controls whether @@ -34,7 +29,6 @@ class SpecDecodePlan: dynamic_batch_size: Optional[int] draft_step: int pre_draft_step: int - selection_mode: str = "confidence" @property def is_dynamic(self) -> bool: @@ -45,41 +39,36 @@ def skip_verify_sync(self) -> bool: return self.is_dynamic and self.pre_draft_step == 0 -class FixedMTPPlanner: - """Planner for static MTP.""" +class FixedSpecPlanner: + """Planner for fixed-width speculative decoding.""" - def __init__(self, mtp_step: int) -> None: - self.mtp_step = int(mtp_step) + def __init__(self, max_draft_step: int) -> None: + self.max_draft_step = int(max_draft_step) def plan(self, req_num: int | None = None, original_batch_size: int | None = None) -> SpecDecodePlan: - del req_num - del original_batch_size return SpecDecodePlan( dynamic_batch_size=None, - draft_step=self.mtp_step, - pre_draft_step=self.mtp_step, - selection_mode="none", + draft_step=self.max_draft_step, + pre_draft_step=self.max_draft_step, ) -class DynamicMTPPlanner: - planner_mode = "default" - +class DynamicSpecPlanner: def __init__( self, - mtp_step: int, + max_draft_step: int, use_random_mode: bool = True, random_mode_iter_threshold: int = 100, ) -> None: - self.mtp_step = int(mtp_step) + self.max_draft_step = int(max_draft_step) # 用于记录 decode 时的静态推理耗时(ms)。 self.main_model_speeds = _InferCostMsTable() self.draft_model_speeds = _InferCostMsTable() - # 记录每个对应长度mtp step 步的接受概率。 由于原始位置必然是接受的,所以不需要记录。 - self.mtp_len_to_accept_ratio = [ - _EMAValue(decay=0.95, init_value=1.0, enable_decay_warmup=False) for _ in range(self.mtp_step) + # 记录不同 draft 深度的接受概率;第一个 target token 必然接受,不需要统计。 + self.draft_len_to_accept_ratio = [ + _EMAValue(decay=0.95, init_value=1.0, enable_decay_warmup=False) for _ in range(self.max_draft_step) ] # 记录请求数量以及对应的推理dynamic_batch_size 对应的接受率统计 self.req_num_to_dynamic_batch_size_to_accept_ratio: Dict[int, Dict[int, _EMAValue]] = {} @@ -91,52 +80,38 @@ def __init__( self._random = random.Random(0) # 记录上一次选择的draft step 步长,才好选择对应的 dynamic_batch_size - self.pre_draft_step = self.mtp_step - self._selection_mode = "confidence" - return + self.pre_draft_step = self.max_draft_step def plan(self, req_num: int, original_batch_size: int) -> SpecDecodePlan: dynamic_batch_size, draft_step, pre_draft_step = self.get_dynamic_batch_size( req_num=req_num, original_batch_size=original_batch_size, ) - if dynamic_batch_size == original_batch_size and self._selection_mode == "confidence": - # Full-width dynamic plans do not need confidence sampling or - # tensor compaction, but ordinary calibration/controller plans - # still need to produce confidence for a potentially narrower - # next iteration. ``observe`` skips compaction while retaining - # those probabilities; the explicit profitable-chain path below - # uses ``full`` for the true Static-equivalent fast path. - self._selection_mode = "observe" return SpecDecodePlan( dynamic_batch_size=dynamic_batch_size, draft_step=draft_step, pre_draft_step=pre_draft_step, - selection_mode=self._selection_mode, ) - def update_infer_cost(self, *, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: + def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: speed_table = self.draft_model_speeds if is_draft_model else self.main_model_speeds speed_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - return - def update_mtp_len_to_accept_ratio(self, mtp_len: int, accept_ratio: float) -> None: - assert mtp_len > 0 and mtp_len <= self.mtp_step - self.mtp_len_to_accept_ratio[mtp_len - 1].update(accept_ratio) - return + def update_draft_len_to_accept_ratio(self, draft_len: int, accept_ratio: float) -> None: + assert draft_len > 0 and draft_len <= self.max_draft_step + self.draft_len_to_accept_ratio[draft_len - 1].update(accept_ratio) - def update_verified_prefix_stats(self, *, verify_len: int, accept_len: int) -> None: + def update_verified_prefix_stats(self, verify_len: int, accept_len: int) -> None: if verify_len - 1 <= 0: return - for mtp_index in range(verify_len - 1): - mtp_len = mtp_index + 1 - ratio = (accept_len - 1) / mtp_len + for draft_index in range(verify_len - 1): + draft_len = draft_index + 1 + ratio = (accept_len - 1) / draft_len ratio = max(0.0, min(1.0, ratio)) - self.update_mtp_len_to_accept_ratio( - mtp_len=mtp_len, + self.update_draft_len_to_accept_ratio( + draft_len=draft_len, accept_ratio=ratio, ) - return def update_req_num_to_dynamic_batch_size_to_accept_ratio( self, req_num: int, dynamic_batch_size: int, accept_ratio: float @@ -145,7 +120,6 @@ def update_req_num_to_dynamic_batch_size_to_accept_ratio( self._get_req_num_to_dynamic_batch_size_to_accept_ratio( req_num=req_num, dynamic_batch_size=dynamic_batch_size ).update(accept_ratio) - return def _get_req_num_to_dynamic_batch_size_to_accept_ratio(self, req_num: int, dynamic_batch_size: int) -> "_EMAValue": assert dynamic_batch_size >= req_num @@ -164,19 +138,19 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl 调用方可据此判断当前 verify 是否有真实候选需要验证 (pre_draft_step == 0 时 accept_len 恒为 1,无需等待 GPU verify 结果)。 """ - assert req_num * (self.mtp_step + 1) == original_batch_size + assert req_num * (self.max_draft_step + 1) == original_batch_size pre_draft_step = self.pre_draft_step if req_num == 0: - self.pre_draft_step = self.mtp_step - return 0, self.mtp_step, pre_draft_step + self.pre_draft_step = self.max_draft_step + return 0, self.max_draft_step, pre_draft_step if not self.main_model_speeds.has_data() or not self.draft_model_speeds.has_data(): # The cost model is only meaningful after both target and draft # decode costs have been profiled. Block proposers such as DFlash # do not run through draft_model.forward, and cudagraph may also be - # disabled, so a missing table must not collapse dynamic MTP to + # disabled, so a missing table must not collapse dynamic scheduling to # draft_step=0. - self.pre_draft_step = self.mtp_step - return req_num * (pre_draft_step + 1), self.mtp_step, pre_draft_step + self.pre_draft_step = self.max_draft_step + return req_num * (pre_draft_step + 1), self.max_draft_step, pre_draft_step # case 1 如果采用随机的方式决定 dynamic_batch_size self._iter += 1 @@ -185,7 +159,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl max_batch_size = req_num * (pre_draft_step + 1) dynamic_batch_size = self._random.randint(min_batch_size, max_batch_size) - draft_step = self._random.randint(0, self.mtp_step) + draft_step = self._random.randint(0, self.max_draft_step) self.pre_draft_step = draft_step return dynamic_batch_size, draft_step, pre_draft_step @@ -204,14 +178,14 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl # 下一步的 draft step 选择,需要考虑计算不同step步的收益问题再决定 min_cost_ms = float("inf") min_cost_ms_draft_step = 0 # 默认选择0步长 - for draft_step in range(0, self.mtp_step + 1): + for draft_step in range(0, self.max_draft_step + 1): cost_ms = self._get_cost_ms(req_num=req_num, dynamic_batch_size=dynamic_batch_size, draft_step=draft_step) if cost_ms < min_cost_ms: min_cost_ms = cost_ms min_cost_ms_draft_step = draft_step - # draft step 步长不能超过 mtp_step, 也不能小于0 - min_cost_ms_draft_step = min(min_cost_ms_draft_step, self.mtp_step) + # draft step 步长不能超过 max_draft_step, 也不能小于0 + min_cost_ms_draft_step = min(min_cost_ms_draft_step, self.max_draft_step) min_cost_ms_draft_step = max(min_cost_ms_draft_step, 0) self.pre_draft_step = min_cost_ms_draft_step return dynamic_batch_size, min_cost_ms_draft_step, pre_draft_step @@ -242,18 +216,18 @@ def _get_dynamic_batch_size_to_accept_ratio(self, req_num: int, dynamic_batch_si assert real_step >= 1.0 real_step = real_step - 1.0 - # 用插值的方式估计不同mtp_len 对应的接受率 + # 用插值的方式估计不同draft_len 对应的接受率 left = int(math.floor(real_step)) right = int(left + 1) if left == 0: left_value = 0.0 else: - left_value = self.mtp_len_to_accept_ratio[left - 1].get() + left_value = self.draft_len_to_accept_ratio[left - 1].get() - if right > self.mtp_step: + if right > self.max_draft_step: right_value = 0.0 else: - right_value = self.mtp_len_to_accept_ratio[right - 1].get() + right_value = self.draft_len_to_accept_ratio[right - 1].get() accept_ratio = left_value + (right_value - left_value) * (real_step - left) calcu_accept_ratio = (req_num + (dynamic_batch_size - req_num) * accept_ratio) / dynamic_batch_size @@ -262,7 +236,7 @@ def _get_dynamic_batch_size_to_accept_ratio(self, req_num: int, dynamic_batch_si return calcu_accept_ratio * (1 - weight) + ema.get() * weight -class Eagle3DynamicMTPPlanner(DynamicMTPPlanner): +class Eagle3DynamicSpecPlanner(DynamicSpecPlanner): """Joint draft-length and verify-capacity planner for Eagle3. ``pre_draft_step`` bounds the proposal that is being verified now, while @@ -277,19 +251,18 @@ class Eagle3DynamicMTPPlanner(DynamicMTPPlanner): one accepted tail per request (batch B). """ - planner_mode = "eagle3" _ACCEPT_RATIO_BUCKETS_PER_DRAFT_ROW = 8 - def __init__(self, mtp_step: int) -> None: - super().__init__(mtp_step=mtp_step, use_random_mode=False) + def __init__(self, max_draft_step: int) -> None: + super().__init__(max_draft_step=max_draft_step, use_random_mode=False) # Eagle uses these values as a full-verify survival curve. Start from # the first batch mean instead of decaying slowly from an all-accepted # prior; otherwise a 32-iteration calibration still substantially # overestimates short draft depths. self._prefix_survival_decay = float(os.getenv("LIGHTLLM_EAGLE3_PREFIX_SURVIVAL_DECAY", "0.95")) - self.mtp_len_to_accept_ratio = [ + self.draft_len_to_accept_ratio = [ _EMAValue(decay=self._prefix_survival_decay, init_value=1.0, enable_decay_warmup=True) - for _ in range(self.mtp_step) + for _ in range(self.max_draft_step) ] self._min_static_progress_ratio = float(os.getenv("LIGHTLLM_EAGLE3_MIN_STATIC_PROGRESS_RATIO", "0.85")) # This is the externally visible verify efficiency: @@ -315,13 +288,6 @@ def __init__(self, mtp_step: int) -> None: 0, int(os.getenv("LIGHTLLM_EAGLE3_FULL_VERIFY_INTERVAL", "128")), ) - self._early_full_probe_accept_ratio = float( - os.getenv("LIGHTLLM_EAGLE3_EARLY_FULL_PROBE_ACCEPT_RATIO", "0.72") - ) - self._early_full_probe_interval = max( - 1, - int(os.getenv("LIGHTLLM_EAGLE3_EARLY_FULL_PROBE_INTERVAL", "32")), - ) self._progress_relax_ratio = float(os.getenv("LIGHTLLM_EAGLE3_PROGRESS_RELAX_RATIO", "1.0")) self._capacity_accept_ratio_floor = float(os.getenv("LIGHTLLM_EAGLE3_CAPACITY_ACCEPT_RATIO_FLOOR", "0.80")) self._capacity_feedback_gain = float(os.getenv("LIGHTLLM_EAGLE3_CAPACITY_FEEDBACK_GAIN", "0.10")) @@ -338,53 +304,10 @@ def __init__(self, mtp_step: int) -> None: self._max_dynamic_draft_step = max( 0, min( - self.mtp_step, - int(os.getenv("LIGHTLLM_EAGLE3_MAX_DYNAMIC_DRAFT_STEP", str(self.mtp_step))), - ), - ) - self._three_regime_enabled = os.getenv("LIGHTLLM_EAGLE3_THREE_REGIME", "0").lower() in { - "1", - "true", - "yes", - "on", - } - self._adaptive_prefix_enabled = os.getenv("LIGHTLLM_EAGLE3_ADAPTIVE_PREFIX", "0").lower() in { - "1", - "true", - "yes", - "on", - } - self._adaptive_prefix_min_depth = max( - 0, - min( - self.mtp_step, - int(os.getenv("LIGHTLLM_EAGLE3_ADAPTIVE_PREFIX_MIN_DEPTH", "1")), + self.max_draft_step, + int(os.getenv("LIGHTLLM_EAGLE3_MAX_DYNAMIC_DRAFT_STEP", str(self.max_draft_step))), ), ) - self._adaptive_prefix_cost_tolerance = max( - 0.0, - float(os.getenv("LIGHTLLM_EAGLE3_ADAPTIVE_PREFIX_COST_TOLERANCE", "0.02")), - ) - self._adaptive_prefix_long_chain_survival = float( - os.getenv("LIGHTLLM_EAGLE3_ADAPTIVE_PREFIX_LONG_CHAIN_SURVIVAL", "0.55") - ) - self._low_load_max_req_num = max( - 1, - int(os.getenv("LIGHTLLM_EAGLE3_LOW_LOAD_MAX_REQ_NUM", "16")), - ) - self._mid_load_max_req_num = max( - self._low_load_max_req_num, - int(os.getenv("LIGHTLLM_EAGLE3_MID_LOAD_MAX_REQ_NUM", "127")), - ) - self._low_load_prefix_depth = max( - 0, - int(os.getenv("LIGHTLLM_EAGLE3_LOW_LOAD_PREFIX_DEPTH", "7")), - ) - self._mid_load_prefix_depth = max( - 0, - int(os.getenv("LIGHTLLM_EAGLE3_MID_LOAD_PREFIX_DEPTH", "3")), - ) - self._active_req_num = 0 self._full_verify_baseline_decay = float( os.getenv( "LIGHTLLM_EAGLE3_BATCH_ACCEPT_EMA_DECAY", @@ -396,10 +319,8 @@ def __init__(self, mtp_step: int) -> None: assert 0.0 < self._progress_relax_ratio <= 1.0 assert 0.0 < self._capacity_accept_ratio_floor <= 1.0 assert 0.0 < self._capacity_feedback_gain <= 1.0 - assert 0.0 <= self._early_full_probe_accept_ratio <= 1.0 assert 0.0 <= self._prefix_survival_decay < 1.0 assert 0.0 <= self._full_verify_baseline_decay < 1.0 - assert 0.0 <= self._adaptive_prefix_long_chain_survival <= 1.0 # Exact (B, K) statistics are sparse because the live request batch B # changes constantly. Pool observations by normalized selected draft @@ -413,24 +334,18 @@ def __init__(self, mtp_step: int) -> None: # a static baseline. self._full_verify_tokens_per_req_ema = _EMAValue( decay=self._full_verify_baseline_decay, - init_value=float(self.mtp_step + 1), + init_value=float(self.max_draft_step + 1), enable_decay_warmup=True, ) - self._full_verify_tokens_per_req_value = float(self.mtp_step + 1) + self._full_verify_tokens_per_req_value = float(self.max_draft_step + 1) self._full_verify_accepted_token_sum = 0.0 self._full_verify_request_count = 0 self._full_verify_update_count = 0 self._full_probe_pending = False - self._last_full_probe_plan_count = -self._early_full_probe_interval # This feedback comes from real non-full dynamic iterations. It # provides a conservative capacity floor when a sparse cost-table # candidate has an over-optimistic expected-token estimate. - self._observed_dynamic_draft_accept_ratio = _EMAValue( - decay=0.9, - init_value=0.75, - enable_decay_warmup=True, - ) self._observed_dynamic_project_accept_ratio = _EMAValue( decay=0.9, init_value=self._min_project_accept_ratio, @@ -453,13 +368,9 @@ def __init__(self, mtp_step: int) -> None: # toward the largest width that still satisfies project acceptance. self._target_verify_rows_per_req_value: Optional[float] = None - self._plan_log_interval = max(0, int(os.getenv("LIGHTLLM_EAGLE3_PLAN_LOG_INTERVAL", "0"))) self._plan_count = 0 - self._draft_step_counts = Counter() - self._verify_rows_per_req_sum = 0.0 - self._expected_tokens_per_req_sum = 0.0 - def update_verified_prefix_stats(self, *, verify_len: int, accept_len: int) -> None: + def update_verified_prefix_stats(self, verify_len: int, accept_len: int) -> None: """Record the survival probability of each Eagle draft position. The generic planner records ``(accept_len - 1) / depth``. That value @@ -468,31 +379,30 @@ def update_verified_prefix_stats(self, *, verify_len: int, accept_len: int) -> N probability that a token at each depth is actually reached. """ - max_mtp_len = min(max(0, verify_len - 1), self.mtp_step) - for mtp_len in range(1, max_mtp_len + 1): - self.update_mtp_len_to_accept_ratio( - mtp_len=mtp_len, - accept_ratio=1.0 if accept_len > mtp_len else 0.0, + max_draft_len = min(max(0, verify_len - 1), self.max_draft_step) + for draft_len in range(1, max_draft_len + 1): + self.update_draft_len_to_accept_ratio( + draft_len=draft_len, + accept_ratio=1.0 if accept_len > draft_len else 0.0, ) def update_verified_batch_prefix_stats( self, - *, verify_and_accept_lengths: List[Tuple[int, int]], ) -> None: """Update each depth once with a request-weighted batch mean.""" - for mtp_len in range(1, self.mtp_step + 1): + for draft_len in range(1, self.max_draft_step + 1): eligible_accept_lengths = [ - accept_len for verify_len, accept_len in verify_and_accept_lengths if verify_len > mtp_len + accept_len for verify_len, accept_len in verify_and_accept_lengths if verify_len > draft_len ] if not eligible_accept_lengths: continue - survival_ratio = sum(accept_len > mtp_len for accept_len in eligible_accept_lengths) / len( + survival_ratio = sum(accept_len > draft_len for accept_len in eligible_accept_lengths) / len( eligible_accept_lengths ) - self.update_mtp_len_to_accept_ratio( - mtp_len=mtp_len, + self.update_draft_len_to_accept_ratio( + draft_len=draft_len, accept_ratio=survival_ratio, ) @@ -501,9 +411,9 @@ def update_req_num_to_dynamic_batch_size_to_accept_ratio( req_num: int, dynamic_batch_size: int, accept_ratio: float, - verify_step: int = None, + pre_draft_step: int = None, ) -> None: - depth = self.mtp_step if verify_step is None else int(verify_step) + depth = self.max_draft_step if pre_draft_step is None else int(pre_draft_step) exact_key = (depth, int(req_num), int(dynamic_batch_size)) if exact_key not in self._accept_ratio_by_depth_req_and_batch: self._accept_ratio_by_depth_req_and_batch[exact_key] = _EMAValue( @@ -515,11 +425,11 @@ def update_req_num_to_dynamic_batch_size_to_accept_ratio( self._get_width_bucket_accept_ratio( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - verify_step=verify_step, + pre_draft_step=pre_draft_step, ).update(accept_ratio) def update_full_verify_tokens_per_req(self, tokens_per_req: float, req_num: int = 1) -> None: - tokens_per_req = max(1.0, min(float(self.mtp_step + 1), float(tokens_per_req))) + tokens_per_req = max(1.0, min(float(self.max_draft_step + 1), float(tokens_per_req))) req_num = max(1, int(req_num)) self._full_verify_accepted_token_sum += tokens_per_req * req_num self._full_verify_request_count += req_num @@ -534,7 +444,6 @@ def update_full_verify_tokens_per_req(self, tokens_per_req: float, req_num: int def update_observed_iteration_stats( self, - *, tokens_per_req: float, verify_rows_per_req: float, is_full_verify: bool, @@ -543,10 +452,7 @@ def update_observed_iteration_stats( if is_full_verify or verify_rows_per_req <= 1.0: return req_num = max(1, int(req_num)) - draft_accept_ratio = (tokens_per_req - 1.0) / (verify_rows_per_req - 1.0) - draft_accept_ratio = max(0.0, min(1.0, draft_accept_ratio)) project_accept_ratio = max(0.0, min(1.0, tokens_per_req / verify_rows_per_req)) - self._observed_dynamic_draft_accept_ratio.update(draft_accept_ratio) self._observed_dynamic_project_accept_ratio.update(project_accept_ratio) self._observed_dynamic_tokens_per_req.update(tokens_per_req) self._observed_dynamic_accepted_token_sum += tokens_per_req * req_num @@ -561,7 +467,6 @@ def update_observed_iteration_stats( def _update_target_verify_rows_per_req( self, - *, tokens_per_req: float, verify_rows_per_req: float, req_num: int, @@ -583,7 +488,7 @@ def _update_target_verify_rows_per_req( verify_rows_per_req * min_expected_tokens / max(tokens_per_req, 1e-6), ) - sample_target = max(1.0, min(float(self.mtp_step + 1), sample_target)) + sample_target = max(1.0, min(float(self.max_draft_step + 1), sample_target)) # Bound a single observation. This is especially important while a # batch drains and only a handful of unusually hard requests remain. sample_target = max(current_target - 0.5, min(current_target + 0.5, sample_target)) @@ -596,13 +501,6 @@ def _get_observed_project_accept_ratio(self) -> float: return self._min_project_accept_ratio return self._observed_dynamic_accepted_token_sum / self._observed_dynamic_verify_row_sum - def _get_observed_draft_accept_ratio(self) -> float: - conditional_rows = self._observed_dynamic_verify_row_sum - self._observed_dynamic_request_count - if conditional_rows <= 0.0: - return self._observed_dynamic_draft_accept_ratio.get() - conditional_tokens = self._observed_dynamic_accepted_token_sum - self._observed_dynamic_request_count - return max(0.0, min(1.0, conditional_tokens / conditional_rows)) - def _get_observed_tokens_per_req(self) -> float: if self._observed_dynamic_request_count <= 0: return self._observed_dynamic_tokens_per_req.get() @@ -612,9 +510,9 @@ def _get_dynamic_batch_size_to_accept_ratio( self, req_num: int, dynamic_batch_size: int, - verify_step: int = None, + pre_draft_step: int = None, ): - depth = self.mtp_step if verify_step is None else int(verify_step) + depth = self.max_draft_step if pre_draft_step is None else int(pre_draft_step) exact_key = (depth, int(req_num), int(dynamic_batch_size)) exact_ema = self._accept_ratio_by_depth_req_and_batch.get(exact_key) base_estimate = super()._get_dynamic_batch_size_to_accept_ratio( @@ -627,21 +525,20 @@ def _get_dynamic_batch_size_to_accept_ratio( width_ema = self._get_width_bucket_accept_ratio( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - verify_step=verify_step, + pre_draft_step=pre_draft_step, ) width_weight = min(1.0, width_ema.get_count() / 10.0) return base_estimate * (1.0 - width_weight) + width_ema.get() * width_weight def _get_width_bucket_accept_ratio( self, - *, req_num: int, dynamic_batch_size: int, - verify_step: int = None, + pre_draft_step: int = None, ) -> "_EMAValue": selected_draft_rows_per_req = max(0.0, dynamic_batch_size / req_num - 1.0) width_bucket = int(round(selected_draft_rows_per_req * self._ACCEPT_RATIO_BUCKETS_PER_DRAFT_ROW)) - depth = self.mtp_step if verify_step is None else int(verify_step) + depth = self.max_draft_step if pre_draft_step is None else int(pre_draft_step) bucket = (depth, width_bucket) if bucket not in self._accept_ratio_by_depth_and_width_bucket: self._accept_ratio_by_depth_and_width_bucket[bucket] = _EMAValue( @@ -652,170 +549,34 @@ def _get_width_bucket_accept_ratio( return self._accept_ratio_by_depth_and_width_bucket[bucket] def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: - assert req_num * (self.mtp_step + 1) == original_batch_size - previous_req_num = self._active_req_num - self._active_req_num = int(req_num) - self._selection_mode = "confidence" + assert req_num * (self.max_draft_step + 1) == original_batch_size pre_draft_step = self.pre_draft_step if req_num == 0: - self.pre_draft_step = self.mtp_step - return 0, self.mtp_step, pre_draft_step + self.pre_draft_step = self.max_draft_step + return 0, self.max_draft_step, pre_draft_step max_batch_size = req_num * (pre_draft_step + 1) if not self.main_model_speeds.has_data() or not self.draft_model_speeds.has_data(): - self._selection_mode = "observe" - self.pre_draft_step = self.mtp_step - return max_batch_size, self.mtp_step, pre_draft_step - - # Crossing from the prefix-friendly regime into high load invalidates - # the old per-request capacity target. Carrying that wide low-load - # target into a suddenly larger batch creates several expensive - # iterations before feedback contracts it. Clamp immediately to the - # high-load progress floor; normal closed-loop feedback can widen it - # again if the additional rows remain profitable. - crossed_into_high_load = ( - self._three_regime_enabled - and previous_req_num > 0 - and previous_req_num <= self._mid_load_max_req_num - and req_num > self._mid_load_max_req_num - ) - if crossed_into_high_load: - high_load_target = self._get_target_verify_rows_per_req() - if self._target_verify_rows_per_req_value is None: - self._target_verify_rows_per_req_value = high_load_target - else: - self._target_verify_rows_per_req_value = min( - self._target_verify_rows_per_req_value, - high_load_target, - ) - - # The legacy three-regime policy uses a fixed prefix depth. The - # adaptive variant below delays this decision until after full-probe - # handling, then chooses the prefix from the live survival curve and - # target+draft cost model. - prefix_depth = None if self._adaptive_prefix_enabled else self._get_load_prefix_depth(req_num=req_num) - if prefix_depth is not None: - self._selection_mode = "prefix" - verify_step = min(pre_draft_step, prefix_depth) - draft_step = min(self.mtp_step, prefix_depth) - dynamic_batch_size = req_num * (verify_step + 1) - self.pre_draft_step = draft_step - expected_tokens_per_req = ( - self._estimate_expected_token_num( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - verify_step=pre_draft_step, - ) - / req_num - ) - self._record_plan_stats( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - expected_tokens_per_req=expected_tokens_per_req, - ) - return dynamic_batch_size, draft_step, pre_draft_step + self.pre_draft_step = self.max_draft_step + return max_batch_size, self.max_draft_step, pre_draft_step if self._should_schedule_full_probe(): self._full_probe_pending = True force_full_verify = self._should_force_full_verify(pre_draft_step=pre_draft_step) if force_full_verify: - self._selection_mode = "observe" # Keep drafting full width during initial calibration. A periodic # probe may return to the normal dynamic draft length immediately # after its one full target verify. draft_step = ( - self.mtp_step + self.max_draft_step if self._full_verify_update_count < self._full_verify_warmup_steps else self._select_next_draft_step(req_num=req_num) ) dynamic_batch_size = max_batch_size self.pre_draft_step = draft_step - expected_tokens_per_req = ( - self._estimate_expected_token_num( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - verify_step=pre_draft_step, - ) - / req_num - ) - self._record_plan_stats( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - expected_tokens_per_req=expected_tokens_per_req, - ) - return dynamic_batch_size, draft_step, pre_draft_step - - # A highly predictable workload does not benefit from spending GPU - # work on confidence selection and row compaction. Keep every row of - # the current proposal and restore/retain a K-wide next proposal. On - # the following iteration ``DynamicMTPPlanner.plan`` marks the full - # width as a runtime fast path, so its target forward is identical to - # Static EAGLE3 while the planner continues monitoring acceptance. - # This rule is intentionally load-independent: high-concurrency, - # high-acceptance batches should expand just like low-load ones. - if ( - self._adaptive_prefix_enabled - and req_num <= self._mid_load_max_req_num - and self._has_profitable_long_chain() - ): - # If the previous proposal was shortened, only that prefix is - # valid even though ModelInput remains padded to the static width. - # Verify the valid prefix once while rebuilding K; the following - # iteration can then enter the true full-width fast path. - self._selection_mode = "full" if pre_draft_step == self.mtp_step else "prefix" - dynamic_batch_size = max_batch_size - draft_step = self.mtp_step - self.pre_draft_step = draft_step - expected_tokens_per_req = ( - self._estimate_expected_token_num( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - verify_step=pre_draft_step, - ) - / req_num - ) - self._record_plan_stats( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - expected_tokens_per_req=expected_tokens_per_req, - ) - return dynamic_batch_size, draft_step, pre_draft_step - - adaptive_prefix_limit = self._get_adaptive_prefix_limit(req_num=req_num) - if adaptive_prefix_limit is not None: - self._selection_mode = "prefix" - prefix_depth = self._select_adaptive_prefix_depth( - req_num=req_num, - max_depth=adaptive_prefix_limit, - ) - verify_step = min(pre_draft_step, prefix_depth) - dynamic_batch_size = req_num * (verify_step + 1) - # A full probe needs two iterations after the proposal has been - # shortened: first rebuild a K-wide proposal, then verify it on - # the next target forward. Without this preparation step the - # pending probe can never observe deep positions and the planner - # becomes unable to recover a long chain after a workload shift. - draft_step = self.mtp_step if self._full_probe_pending else prefix_depth - self.pre_draft_step = draft_step - expected_tokens_per_req = ( - self._estimate_expected_token_num( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - verify_step=pre_draft_step, - ) - / req_num - ) - self._record_plan_stats( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - expected_tokens_per_req=expected_tokens_per_req, - ) + self._record_plan() return dynamic_batch_size, draft_step, pre_draft_step # Once the full-width baseline is calibrated, choose draft depth by @@ -826,7 +587,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl if self._full_verify_request_count > 0: draft_step = self._select_next_draft_step(req_num=req_num) if self._full_probe_pending: - draft_step = self.mtp_step + draft_step = self.max_draft_step target_verify_rows_per_req = self._get_controlled_verify_rows_per_req() dynamic_batch_size = int(math.ceil(req_num * target_verify_rows_per_req)) if self._align_verify_rows_to_graph: @@ -842,20 +603,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl dynamic_batch_size = graph_batch_size dynamic_batch_size = min(max(dynamic_batch_size, req_num), max_batch_size) self.pre_draft_step = draft_step - expected_tokens_per_req = ( - self._estimate_expected_token_num( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - verify_step=pre_draft_step, - ) - / req_num - ) - self._record_plan_stats( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - expected_tokens_per_req=expected_tokens_per_req, - ) + self._record_plan() return dynamic_batch_size, draft_step, pre_draft_step # The action selected here builds the proposal consumed by the next @@ -863,7 +611,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl # width so a zero-step iteration can recover on its own. draft_step = self._select_next_draft_step(req_num=req_num) if self._full_probe_pending: - draft_step = self.mtp_step + draft_step = self.max_draft_step dynamic_batch_size_keys = self._get_candidate_batch_sizes( req_num=req_num, @@ -873,21 +621,15 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl self._get_eagle3_candidate( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - verify_step=pre_draft_step, + pre_draft_step=pre_draft_step, draft_step=draft_step, ) for dynamic_batch_size in dynamic_batch_size_keys ] best_candidate = self._select_best_candidate(candidates, req_num=req_num) dynamic_batch_size = best_candidate[1] - expected_tokens_per_req = best_candidate[2] self.pre_draft_step = draft_step - self._record_plan_stats( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - expected_tokens_per_req=expected_tokens_per_req, - ) + self._record_plan() return dynamic_batch_size, draft_step, pre_draft_step def _get_target_verify_rows_per_req(self) -> float: @@ -895,95 +637,19 @@ def _get_target_verify_rows_per_req(self) -> float: if min_expected_tokens <= 1.0: return 1.0 return min( - float(self.mtp_step + 1), + float(self.max_draft_step + 1), min_expected_tokens / self._min_project_accept_ratio, ) - def _get_load_prefix_depth(self, *, req_num: int) -> Optional[int]: - if not self._three_regime_enabled: - return None - if req_num <= self._low_load_max_req_num: - return min(self.mtp_step, self._low_load_prefix_depth) - if req_num <= self._mid_load_max_req_num: - return min(self.mtp_step, self._mid_load_prefix_depth) - return None - - def _get_adaptive_prefix_limit(self, *, req_num: int) -> Optional[int]: - if not self._adaptive_prefix_enabled or not self._three_regime_enabled: - return None - if req_num <= self._low_load_max_req_num: - return min(self.mtp_step, self._low_load_prefix_depth) - if req_num <= self._mid_load_max_req_num: - return min(self.mtp_step, self._mid_load_prefix_depth) - return None - - def _select_adaptive_prefix_depth(self, *, req_num: int, max_depth: int) -> int: - """Choose a steady-state chain depth from measured prefix survival. - - Prefix selection is appropriate at low load because preserving a - request's chain can turn spare target capacity into forward progress. - The expected accepted length is computed directly from the survival - probability of every draft position, while the cost includes both the - target verification and Eagle3 proposal construction. If several - depths are effectively tied, prefer the longer chain so predictable - requests can retain near-K accepted lengths. - """ - - max_depth = max(self._adaptive_prefix_min_depth, min(self.mtp_step, int(max_depth))) - min_depth = min(self._adaptive_prefix_min_depth, max_depth) - if self._has_profitable_long_chain(max_depth=max_depth): - return max_depth - - candidates = [] - for depth in range(min_depth, max_depth + 1): - expected_tokens_per_req = 1.0 + sum( - self.mtp_len_to_accept_ratio[mtp_index].get() for mtp_index in range(depth) - ) - dynamic_batch_size = req_num * (depth + 1) - cost_ms = self._get_eagle3_cost_ms( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - verify_step=depth, - draft_step=depth, - expected_token_num=req_num * expected_tokens_per_req, - ) - project_accept_ratio = expected_tokens_per_req / (depth + 1) - candidates.append((cost_ms, depth, project_accept_ratio)) - - feasible = [ - candidate for candidate in candidates if candidate[2] >= self._min_project_accept_ratio - ] - if not feasible: - best_accept_ratio = max(candidate[2] for candidate in candidates) - feasible = [ - candidate for candidate in candidates if candidate[2] >= best_accept_ratio - 1e-6 - ] - - best_cost = min(cost for cost, _, _ in feasible) - return max( - depth - for cost, depth, _ in feasible - if cost <= best_cost * (1.0 + self._adaptive_prefix_cost_tolerance) - ) - - def _has_profitable_long_chain(self, *, max_depth: int = None) -> bool: - max_depth = self.mtp_step if max_depth is None else min(self.mtp_step, int(max_depth)) - if max_depth <= 0: - return False - mean_survival = sum( - self.mtp_len_to_accept_ratio[mtp_index].get() for mtp_index in range(max_depth) - ) / max_depth - return mean_survival >= self._adaptive_prefix_long_chain_survival - def _get_controlled_verify_rows_per_req(self) -> float: if self._target_verify_rows_per_req_value is None: return self._get_target_verify_rows_per_req() return max( 1.0, - min(float(self.mtp_step + 1), self._target_verify_rows_per_req_value), + min(float(self.max_draft_step + 1), self._target_verify_rows_per_req_value), ) - def _get_candidate_batch_sizes(self, *, req_num: int, max_batch_size: int) -> List[int]: + def _get_candidate_batch_sizes(self, req_num: int, max_batch_size: int) -> List[int]: candidates = set( self.main_model_speeds.get_batch_size_keys_between( req_num, @@ -1010,35 +676,15 @@ def _get_candidate_batch_sizes(self, *, req_num: int, max_batch_size: int) -> Li return sorted(candidates) def _should_schedule_full_probe(self) -> bool: - periodic_probe = ( + return ( self._full_verify_interval > 0 and self._full_verify_update_count >= self._full_verify_warmup_steps and self._plan_count > 0 and self._plan_count % self._full_verify_interval == 0 ) - # At low/mid load, confidence-selected prefix rows expose a rising - # acceptance regime before the periodic unbiased K-wide probe arrives. - # Probe early once prefix acceptance is high enough, so the survival - # curve can discover newly profitable deep tokens and restore long - # chains without waiting tens of seconds. Keep a cooldown and disable - # this path at high load, where a full probe is materially expensive. - early_probe = ( - self._adaptive_prefix_enabled - and self._active_req_num > 0 - and self._active_req_num <= self._mid_load_max_req_num - and self._observed_dynamic_draft_accept_ratio.get() - >= self._early_full_probe_accept_ratio - and self._plan_count - self._last_full_probe_plan_count - >= self._early_full_probe_interval - and not self._has_profitable_long_chain() - ) - if periodic_probe or early_probe: - self._last_full_probe_plan_count = self._plan_count - return True - return False - def _should_force_full_verify(self, *, pre_draft_step: int) -> bool: - if pre_draft_step != self.mtp_step: + def _should_force_full_verify(self, pre_draft_step: int) -> bool: + if pre_draft_step != self.max_draft_step: return False if self._full_verify_update_count < self._full_verify_warmup_steps: return True @@ -1047,55 +693,10 @@ def _should_force_full_verify(self, *, pre_draft_step: int) -> bool: return True return False - def _record_plan_stats( - self, - *, - req_num: int, - dynamic_batch_size: int, - draft_step: int, - expected_tokens_per_req: float, - ) -> None: + def _record_plan(self) -> None: self._plan_count += 1 - if self._plan_log_interval <= 0: - return - self._draft_step_counts[int(draft_step)] += 1 - self._verify_rows_per_req_sum += dynamic_batch_size / req_num - self._expected_tokens_per_req_sum += expected_tokens_per_req - if self._plan_count % self._plan_log_interval == 0: - logger.info( - "eagle3_dynamic_plan_stats plan_count=%d draft_step_counts=%s " - "avg_verify_rows_per_req=%.6f avg_expected_tokens_per_req=%.6f " - "static_tokens_per_req=%.6f full_verify_count=%d full_verify_req_count=%d " - "observed_project_accept_ratio=%.6f observed_draft_accept_ratio=%.6f " - "observed_tokens_per_req=%.6f target_verify_rows_per_req=%.6f " - "prefix_accept_ratios=%s", - self._plan_count, - dict(sorted(self._draft_step_counts.items())), - self._verify_rows_per_req_sum / self._plan_count, - self._expected_tokens_per_req_sum / self._plan_count, - self._full_verify_tokens_per_req_value, - self._full_verify_update_count, - self._full_verify_request_count, - self._get_observed_project_accept_ratio(), - self._get_observed_draft_accept_ratio(), - self._get_observed_tokens_per_req(), - self._get_controlled_verify_rows_per_req(), - [round(value.get(), 6) for value in self.mtp_len_to_accept_ratio], - ) - def get_trace_stats(self) -> dict: - """Expose controller state for controlled scheduling ablations.""" - - return { - "batch_accept_len_estimate": self._full_verify_tokens_per_req_value, - "batch_accept_ema_decay": self._full_verify_baseline_decay, - "batch_accept_full_verify_updates": self._full_verify_update_count, - "target_verify_rows_per_req": self._get_controlled_verify_rows_per_req(), - "observed_tokens_per_req": self._get_observed_tokens_per_req(), - "min_expected_tokens_per_req": self._get_min_expected_tokens_per_req(), - } - - def _select_next_draft_step(self, *, req_num: int) -> int: + def _select_next_draft_step(self, req_num: int) -> int: """Choose a recoverable long-run Eagle3 draft length. For each possible length, jointly search the verify capacities that @@ -1112,7 +713,7 @@ def _select_next_draft_step(self, *, req_num: int) -> int: # optimistic sparse (B, K) bucket from selecting draft_step=3 on # GSM8K and paying for many extra target iterations. max_depth_tokens_per_req = 1.0 + sum( - self.mtp_len_to_accept_ratio[mtp_index].get() for mtp_index in range(draft_step) + self.draft_len_to_accept_ratio[draft_index].get() for draft_index in range(draft_step) ) if max_depth_tokens_per_req + 1e-6 < min_expected_tokens: continue @@ -1126,7 +727,7 @@ def _select_next_draft_step(self, *, req_num: int) -> int: self._get_eagle3_candidate( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - verify_step=draft_step, + pre_draft_step=draft_step, draft_step=draft_step, ) ) @@ -1136,21 +737,20 @@ def _select_next_draft_step(self, *, req_num: int) -> int: def _get_eagle3_candidate( self, - *, req_num: int, dynamic_batch_size: int, - verify_step: int, + pre_draft_step: int, draft_step: int, ) -> Tuple[float, int, float, float, int]: expected_token_num = self._estimate_expected_token_num( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - verify_step=verify_step, + pre_draft_step=pre_draft_step, ) cost_ms = self._get_eagle3_cost_ms( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - verify_step=verify_step, + pre_draft_step=pre_draft_step, draft_step=draft_step, expected_token_num=expected_token_num, ) @@ -1167,7 +767,6 @@ def _get_eagle3_candidate( def _select_best_candidate( self, candidates: List[Tuple[float, int, float, float, int]], - *, req_num: int, ) -> Tuple[float, int, float, float, int]: assert candidates @@ -1218,13 +817,6 @@ def _select_best_candidate( def _get_min_expected_tokens_per_req(self) -> float: if self._full_verify_request_count == 0: return 1.0 - if self._three_regime_enabled and self._active_req_num > self._mid_load_max_req_num: - # Relax only the speculative gain above the guaranteed base row. - # Scaling the total accepted length can collapse the target below - # one token/request and makes the high-load controller degenerate - # into no-MTP. - static_gain = max(0.0, self._full_verify_tokens_per_req_value - 1.0) - return 1.0 + static_gain * self._min_static_progress_ratio * self._progress_relax_ratio target_tokens = max( 1.0, self._full_verify_tokens_per_req_value * self._min_static_progress_ratio, @@ -1243,24 +835,23 @@ def _get_min_verify_rows_per_req(self) -> float: ) return min_expected_tokens / observed_project_accept_ratio - def _estimate_expected_token_num(self, *, req_num: int, dynamic_batch_size: int, verify_step: int) -> float: + def _estimate_expected_token_num(self, req_num: int, dynamic_batch_size: int, pre_draft_step: int) -> float: accept_ratio = self._get_dynamic_batch_size_to_accept_ratio( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - verify_step=verify_step, + pre_draft_step=pre_draft_step, ) expected_token_num = min( dynamic_batch_size * accept_ratio, - req_num * (verify_step + 1), + req_num * (pre_draft_step + 1), ) return max(float(req_num), expected_token_num) def _get_eagle3_cost_ms( self, - *, req_num: int, dynamic_batch_size: int, - verify_step: int, + pre_draft_step: int, draft_step: int, expected_token_num: float = None, ) -> float: @@ -1268,7 +859,7 @@ def _get_eagle3_cost_ms( expected_token_num = self._estimate_expected_token_num( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - verify_step=verify_step, + pre_draft_step=pre_draft_step, ) # Eagle3 commit runs on all selected verify rows. Every recurrent @@ -1283,33 +874,30 @@ def _get_eagle3_cost_ms( return total_time_ms / expected_token_num -class DSparkDynamicMTPPlanner(DynamicMTPPlanner): +class DSparkDynamicSpecPlanner(DynamicSpecPlanner): """DSpark confidence-scheduled verify-capacity planner.""" - planner_mode = "dspark" - - def __init__(self, mtp_step: int) -> None: + def __init__(self, max_draft_step: int) -> None: super().__init__( - mtp_step=mtp_step, + max_draft_step=max_draft_step, use_random_mode=False, ) - self.mtp_len_to_accept_ratio = [ - _EMAValue(decay=0.95, init_value=1.0, enable_decay_warmup=True) for _ in range(self.mtp_step) + self.draft_len_to_accept_ratio = [ + _EMAValue(decay=0.95, init_value=1.0, enable_decay_warmup=True) for _ in range(self.max_draft_step) ] self._predicted_dynamic_batch_sizes = deque(maxlen=2) - def update_verified_prefix_stats(self, *, verify_len: int, accept_len: int) -> None: + def update_verified_prefix_stats(self, verify_len: int, accept_len: int) -> None: if verify_len - 1 <= 0: return - max_mtp_len = min(verify_len - 1, self.mtp_step) - for mtp_len in range(1, max_mtp_len + 1): - self.update_mtp_len_to_accept_ratio( - mtp_len=mtp_len, - accept_ratio=1.0 if accept_len > mtp_len else 0.0, + max_draft_len = min(verify_len - 1, self.max_draft_step) + for draft_len in range(1, max_draft_len + 1): + self.update_draft_len_to_accept_ratio( + draft_len=draft_len, + accept_ratio=1.0 if accept_len > draft_len else 0.0, ) - return - def update_predicted_schedule_probs(self, *, schedule_probs, req_num: int) -> None: + def update_predicted_schedule_probs(self, schedule_probs, req_num: int) -> None: """Record a confidence-derived future capacity estimate. The current decode step still routes rows by the current probabilities @@ -1327,7 +915,7 @@ def update_predicted_schedule_probs(self, *, schedule_probs, req_num: int) -> No if probs is None or probs.ndim != 2 or probs.shape[1] <= 1: return - draft_probs = probs[:, 1 : self.mtp_step + 1] + draft_probs = probs[:, 1 : self.max_draft_step + 1] if draft_probs.size == 0: return @@ -1342,35 +930,29 @@ def update_predicted_schedule_probs(self, *, schedule_probs, req_num: int) -> No survival_scores=survival_scores, ) self._predicted_dynamic_batch_sizes.append(dynamic_batch_size) - return def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: - assert req_num * (self.mtp_step + 1) == original_batch_size - # ``DynamicMTPPlanner.plan`` rewrites a full-width confidence plan to - # ``observe`` for that iteration. DSpark may select a narrower plan on - # the next iteration, so do not carry the previous iteration's mode - # into the new capacity decision. - self._selection_mode = "confidence" + assert req_num * (self.max_draft_step + 1) == original_batch_size pre_draft_step = self.pre_draft_step - self.pre_draft_step = self.mtp_step + self.pre_draft_step = self.max_draft_step if req_num == 0: - return 0, self.mtp_step, pre_draft_step + return 0, self.max_draft_step, pre_draft_step max_batch_size = req_num * (pre_draft_step + 1) if not self.main_model_speeds.has_data(): - return max_batch_size, self.mtp_step, pre_draft_step + return max_batch_size, self.max_draft_step, pre_draft_step historical_batch_size = self._pop_historical_dynamic_batch_size( req_num=req_num, max_batch_size=max_batch_size, ) if historical_batch_size is not None: - return historical_batch_size, self.mtp_step, pre_draft_step + return historical_batch_size, self.max_draft_step, pre_draft_step if len(self._predicted_dynamic_batch_sizes) > 0: # A confidence estimate is available but has not satisfied the # two-step async delay yet. Keep capacity conservative instead of # leaking a same-step EMA fallback into DSpark scheduling. - return req_num, self.mtp_step, pre_draft_step + return req_num, self.max_draft_step, pre_draft_step candidate_batch_sizes = set(self.main_model_speeds.get_batch_size_keys_between(req_num, max_batch_size)) candidate_batch_sizes.add(req_num) @@ -1392,9 +974,9 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl best_throughput = throughput best_batch_size = dynamic_batch_size - return best_batch_size, self.mtp_step, pre_draft_step + return best_batch_size, self.max_draft_step, pre_draft_step - def _pop_historical_dynamic_batch_size(self, *, req_num: int, max_batch_size: int) -> Optional[int]: + def _pop_historical_dynamic_batch_size(self, req_num: int, max_batch_size: int) -> Optional[int]: if len(self._predicted_dynamic_batch_sizes) < 2: return None predicted_batch_size = int(self._predicted_dynamic_batch_sizes.popleft()) @@ -1402,7 +984,6 @@ def _pop_historical_dynamic_batch_size(self, *, req_num: int, max_batch_size: in def _select_dynamic_batch_size_from_survival_scores( self, - *, req_num: int, survival_scores: np.ndarray, ) -> int: @@ -1440,7 +1021,7 @@ def _select_dynamic_batch_size_from_survival_scores( return best_batch_size @staticmethod - def _topk_prefix_sums(*, values: np.ndarray, counts: List[int]) -> Dict[int, float]: + def _topk_prefix_sums(values: np.ndarray, counts: List[int]) -> Dict[int, float]: """Return sum(top-k(values)) only for the requested k values.""" if len(counts) == 0: @@ -1481,8 +1062,8 @@ def _topk_prefix_sums(*, values: np.ndarray, counts: List[int]) -> Dict[int, flo def _estimate_survival_prefix(self, pre_draft_step: int) -> List[float]: survival_prefix = [1.0] previous = 1.0 - for mtp_index in range(pre_draft_step): - survival = float(self.mtp_len_to_accept_ratio[mtp_index].get()) + for draft_index in range(pre_draft_step): + survival = float(self.draft_len_to_accept_ratio[draft_index].get()) survival = max(0.0, min(previous, survival)) survival_prefix.append(survival) previous = survival @@ -1490,7 +1071,6 @@ def _estimate_survival_prefix(self, pre_draft_step: int) -> List[float]: @staticmethod def _estimate_expected_tokens( - *, req_num: int, dynamic_batch_size: int, survival_prefix: List[float], @@ -1521,7 +1101,6 @@ def __init__(self) -> None: def update(self, batch_size: int, infer_cost_ms: float) -> None: assert batch_size > 0 self.infer_cost_ms_table[int(batch_size)] = float(infer_cost_ms) - return def has_data(self) -> bool: return len(self.infer_cost_ms_table) > 0 @@ -1542,7 +1121,7 @@ def get(self, batch_size: int) -> float: max_batch_size = self.infer_cost_ms_table.peekitem(-1)[0] if batch_size > max_batch_size: max_infer_cost_ms = self.infer_cost_ms_table.peekitem(-1)[1] - # 这里面的 1000.0 意义是尽量使后续的估计,当超过最大graph支持的范围的时候,会直接倾向于关闭mtp功能。 + # 超过最大 graph 范围时使用高惩罚,使调度器倾向于关闭 speculative draft。 return max_infer_cost_ms + (batch_size - max_batch_size) * 1000.0 else: # 找到第一个大于等于 batch_size 的 key,并返回它的 value。 @@ -1559,7 +1138,7 @@ def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> Lis else: return ans - def get_ceil_batch_size(self, batch_size: int, *, max_batch_size: int) -> Optional[int]: + def get_ceil_batch_size(self, batch_size: int, max_batch_size: int) -> Optional[int]: """Return the next recorded graph shape without inventing a key. The cost table can be sparse or empty when CUDA graph is disabled, so @@ -1593,18 +1172,13 @@ def __init__(self, decay: float, init_value: float, enable_decay_warmup: bool = self.current_decay = self.decay self.value = init_value - self.second_moment_value = init_value ** 2 self.update_count = 0 def update(self, new_value: float): self.update_count += 1 self.value = self.current_decay * self.value + (1.0 - self.current_decay) * new_value - self.second_moment_value = self.current_decay * self.second_moment_value + (1.0 - self.current_decay) * ( - new_value ** 2 - ) # 更新 current_decay 的值,使得 current_decay 逐渐逼近 decay 的值 self.current_decay = min(self.decay, (self.decay + self.current_decay) / 2.0 + 0.001) - return def get(self) -> float: return self.value @@ -1612,20 +1186,11 @@ def get(self) -> float: def get_count(self) -> int: return self.update_count - def get_second_moment(self) -> float: - return self.second_moment_value - - def get_variance(self) -> float: - return max(0.0, self.second_moment_value - self.value ** 2) - - def get_sigma(self) -> float: - return math.sqrt(self.get_variance()) - __all__ = [ - "DSparkDynamicMTPPlanner", - "DynamicMTPPlanner", - "Eagle3DynamicMTPPlanner", - "FixedMTPPlanner", + "DSparkDynamicSpecPlanner", + "DynamicSpecPlanner", + "Eagle3DynamicSpecPlanner", + "FixedSpecPlanner", "SpecDecodePlan", ] diff --git a/lightllm/server/router/model_infer/speculative/proposers/__init__.py b/lightllm/server/router/model_infer/speculative/proposers/__init__.py index ab5e6fe7ba..cbab4cc3d1 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/__init__.py +++ b/lightllm/server/router/model_infer/speculative/proposers/__init__.py @@ -1,61 +1,33 @@ -from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal +from typing import TYPE_CHECKING +if TYPE_CHECKING: + from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer -def build_spec_proposer(runtime) -> BaseSpecProposer: - spec_config = runtime.spec_config - if spec_config.is_dspark: + +def build_spec_proposer(engine) -> "BaseSpecProposer": + spec_mode = engine.spec_mode + if spec_mode == "dspark": from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer - return DSparkProposer(runtime) - if spec_config.is_dflash: + return DSparkProposer(engine=engine) + if spec_mode == "dflash": from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer - return DFlashProposer(runtime) - if spec_config.is_eagle3: + return DFlashProposer(engine=engine) + if spec_mode == "eagle3": from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer - return Eagle3Proposer(runtime) - if spec_config.uses_recurrent_draft_model: + return Eagle3Proposer(engine=engine) + if spec_mode in ("eagle_with_att", "eagle_no_att", "qwen3next_eagle"): from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer - return EagleMTPProposer(runtime) + return EagleMTPProposer(engine=engine) from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer - return VanillaMTPProposer(runtime) - - -def __getattr__(name): - if name == "DFlashProposer": - from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer - - return DFlashProposer - if name == "DSparkProposer": - from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer - - return DSparkProposer - if name == "EagleMTPProposer": - from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer - - return EagleMTPProposer - if name == "Eagle3Proposer": - from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer - - return Eagle3Proposer - if name == "VanillaMTPProposer": - from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer - - return VanillaMTPProposer - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + return VanillaMTPProposer(engine=engine) __all__ = [ - "BaseSpecProposer", - "DFlashProposer", - "DSparkProposer", - "EagleMTPProposer", - "Eagle3Proposer", - "SpecProposal", - "VanillaMTPProposer", "build_spec_proposer", ] diff --git a/lightllm/server/router/model_infer/speculative/proposers/base.py b/lightllm/server/router/model_infer/speculative/proposers/base.py index e27c543430..fcd45be8f7 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/speculative/proposers/base.py @@ -8,7 +8,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput if TYPE_CHECKING: - from lightllm.server.router.model_infer.speculative.runtime import SpecRuntime + from lightllm.server.router.model_infer.speculative.engine import SpecEngine @dataclass @@ -21,11 +21,11 @@ class SpecProposal: token_ids: [verify_batch, draft_step + 1] - In static MTP, `draft_step == backend.mtp_step`. In dynamic MTP it may be - shorter and the runtime pads before scatter. + With fixed scheduling, `draft_step == backend.max_draft_step`. Dynamic speculative + scheduling may make it shorter, and the engine pads before scatter. `draft_probs` is intentionally narrower than DeepSpec's full - [B, K, vocab] probability tensor. Current LightLLM dynamic MTP only needs + [B, K, vocab] probability tensor. The current dynamic scheduler only needs the selected-token probability from each draft step: draft_probs[i]: [verify_batch] @@ -33,7 +33,7 @@ class SpecProposal: `schedule_probs` optionally overrides `draft_probs` for dynamic verify row selection. It can be a list of per-step vectors or a dense [verify_batch, draft_step] matrix. DSpark uses confidence-head conditional - acceptance probabilities here; runtime scatters them into the same + acceptance probabilities here; the engine scatters them into the same per-request buffer and the dynamic selector converts them to prefix survival probabilities. @@ -49,9 +49,6 @@ class SpecProposal: extra_mem_indexes_cpu: Optional[torch.Tensor] draft_probs: Optional[List[torch.Tensor]] = None schedule_probs: Optional[Union[List[torch.Tensor], torch.Tensor]] = None - # Actual rows processed by draft-model forwards. Recurrent Eagle can - # prune low-confidence deep chains, so this need not equal B * draft_step. - draft_forward_rows: Optional[int] = None class BaseSpecProposer: @@ -59,54 +56,47 @@ class BaseSpecProposer: A proposer owns the draft-side state transition. The target model gives it the current target token ids plus captured target hidden features through - SpecRuntime.prepare_draft_* methods. The proposer returns candidate ids + SpecEngine.prepare_draft_* methods. The proposer returns candidate ids but does not verify acceptance; verification is handled by SpecVerifier. """ - def __init__(self, runtime: "SpecRuntime") -> None: - self.runtime = runtime - self.backend = runtime.backend + def __init__(self, engine: "SpecEngine") -> None: + self.engine = engine + self.backend = engine.backend @property - def enable_dynamic_mtp(self) -> bool: - return self.runtime.enable_dynamic_mtp + def enable_dynamic_spec(self) -> bool: + return self.engine.enable_dynamic_spec def prepare_draft_prefill_input( self, - *, model_input: ModelInput, next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: Optional[torch.Tensor] = None, - microbatch_index: int = 0, + mtp_draft_input_hiddens: torch.Tensor, ) -> ModelInput: - return self.runtime.prepare_draft_prefill_input( + return self.engine.prepare_draft_prefill_input( model_input=model_input, next_token_ids=next_token_ids, mtp_draft_input_hiddens=mtp_draft_input_hiddens, - microbatch_index=microbatch_index, ) def prepare_draft_decode_input( self, - *, model_input: ModelInput, next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: Optional[torch.Tensor] = None, - microbatch_index: int = 0, + mtp_draft_input_hiddens: torch.Tensor, ) -> ModelInput: - return self.runtime.prepare_draft_decode_input( + return self.engine.prepare_draft_decode_input( model_input=model_input, next_token_ids=next_token_ids, mtp_draft_input_hiddens=mtp_draft_input_hiddens, - microbatch_index=microbatch_index, ) - def select_accepted_tail_rows(self, *, b_req_mtp_start_loc: torch.Tensor, accept_len: torch.Tensor) -> torch.Tensor: + def select_accepted_tail_rows(self, b_req_mtp_start_loc: torch.Tensor, accept_len: torch.Tensor) -> torch.Tensor: return (b_req_mtp_start_loc + accept_len - 1).to(torch.long) def scatter_selected_step_probs( self, - *, selected_rows: torch.Tensor, selected_probs: torch.Tensor, verify_row_count: int, @@ -122,12 +112,12 @@ def scatter_selected_step_probs( def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: """Allocate draft-owned temporary KV slots.""" - return self.runtime.alloc_extra_mem_indexes(token_count) + return self.engine.alloc_extra_mem_indexes(token_count) def build_initial_draft_state( self, - *, model_input: ModelInput, + model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: """Build initial draft KV/state before the first decode verify step. @@ -137,7 +127,7 @@ def build_initial_draft_state( mem_indexes are reused by the draft state builder. - `next_token_ids`: first accepted target token, shape [run_req_num]. - Runtime.prepare_draft_prefill_input injects captured target hidden + SpecEngine.prepare_draft_prefill_input injects captured target hidden features into `mtp_draft_input_hiddens`. This hook only prepares draft-side state. It intentionally does not @@ -150,10 +140,11 @@ def build_initial_draft_state( def build_initial_draft_state_overlap( self, - *, model_input0: ModelInput, + model_output0: ModelOutput, next_token_ids0: torch.Tensor, model_input1: ModelInput, + model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: """Build initial draft state for two overlapped prefill microbatches.""" @@ -162,30 +153,47 @@ def build_initial_draft_state_overlap( def propose_next( self, - *, main_model_input: ModelInput, - main_model_output: Optional[ModelOutput] = None, + main_model_output: Optional[ModelOutput], next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, - verify_result=None, + accept_len: Optional[torch.Tensor] = None, ) -> SpecProposal: """Generate candidate tokens after one target decode forward. Inputs: - - `main_model_input`: target decode ModelInput. In static MTP its - batch is laid out as [req0-main, req0-draft1, ...]. Dynamic MTP may + - `main_model_input`: target decode ModelInput. In the fixed layout its + batch is laid out as [req0-main, req0-draft1, ...]. Dynamic scheduling may compact this batch before target forward. - `next_token_ids`: target sampled ids for rows in `main_model_input`, shape [verify_batch]. - `b_req_mtp_start_loc`: start row for each logical request inside the - MTP-expanded batch, shape [logical_req_num]. + speculative verify batch, shape [logical_req_num]. - `draft_step`: number of candidate draft tokens to produce. - - `verify_result`: optional target verification result from the just - finished target forward. Stateful block proposers use it to commit - the accepted target-hidden segment before preparing the next block. + - `accept_len`: optional accepted-prefix length from the just-finished + target forward. Stateful block proposers use it to commit the + accepted target-hidden segment before preparing the next block. Returns a SpecProposal whose `token_ids[:, 0]` is `next_token_ids`. """ raise NotImplementedError + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor, + draft_step: int, + ) -> SpecProposal: + """Generate a proposal for two DP-overlapped microbatches.""" + + raise NotImplementedError diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index 0393d28344..8b13ca863f 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -5,12 +5,11 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.speculative import BlockDraftLayout from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal class DFlashProposer(BaseSpecProposer): - """DFlash block proposer aligned to the Eagle3 runtime boundary. + """DFlash block proposer aligned to the Eagle3 engine boundary. DFlash remains a non-causal block-prefill draft model, not a recurrent token decoder. The service flow is: @@ -23,17 +22,15 @@ class DFlashProposer(BaseSpecProposer): `SpecProposal.extra_mem_indexes_cpu` """ - variant = "dflash" - @torch.no_grad() def build_initial_draft_state( self, - *, model_input: ModelInput, + model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - del next_token_ids - target_hidden = self.runtime.get_hidden() + target_hidden = model_output.spec_hidden + assert target_hidden is not None if target_hidden.numel() == 0: return @@ -44,35 +41,36 @@ def build_initial_draft_state( # DFlash consumes target hidden states directly on this prefill path. draft_input.mtp_draft_input_hiddens = target_hidden draft_model.forward(draft_input) - return @torch.no_grad() def propose_next( self, - *, main_model_input: ModelInput, - main_model_output: ModelOutput = None, + main_model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, - verify_result=None, + accept_len: torch.Tensor | None = None, ) -> SpecProposal: - del main_model_output - assert 0 <= draft_step <= self.backend.mtp_step - assert verify_result is not None, "DFlash proposal requires target verify result" + assert 0 <= draft_step <= self.backend.max_draft_step + assert accept_len is not None, "DFlash proposal requires target accept lengths" num_reqs = int(b_req_mtp_start_loc.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - layout = self.backend.block_draft_layout - assert block_size == layout.query_block_size - assert verify_result.accept_len.shape[0] == num_reqs + assert accept_len.shape[0] == num_reqs token_ids = next_token_ids.new_full( (next_token_ids.shape[0], draft_step + 1), fill_value=1, ) token_ids[:, 0] = next_token_ids + assert main_model_output is not None and main_model_output.spec_hidden is not None + self.extend_draft_kv_cache( + main_model_input=main_model_input, + target_hidden=main_model_output.spec_hidden, + ) + if draft_step == 0: return SpecProposal( token_ids=token_ids, @@ -80,13 +78,11 @@ def propose_next( draft_probs=None, ) - self.extend_draft_kv_cache(main_model_input=main_model_input) - - # DFlash only drafts from the accepted tail row of each request. Unlike - # MTP, one anchor row expands to a whole non-causal block. + # DFlash drafts from the accepted tail row of each request; one anchor + # row expands to a complete non-causal draft block. selected_rows = self.select_accepted_tail_rows( b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=verify_result.accept_len, + accept_len=accept_len, ) draft_input, draft_mem_indexes_cpu = self.build_block_draft_input( main_model_input=main_model_input, @@ -99,31 +95,19 @@ def propose_next( flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) assert flat_token_ids.numel() == num_reqs * block_size block_token_ids = flat_token_ids.reshape(num_reqs, block_size) - token_ids[selected_rows, 1:] = self.select_draft_token_ids( - block_token_ids=block_token_ids, - draft_step=draft_step, - layout=layout, + # Standard DFlash has one leading bonus row; DeepSpec checkpoints do not. + bonus_rows = block_size - draft_step + assert bonus_rows in (0, 1), ( + f"DFlash block_size={block_size} must equal mtp_step={draft_step} " f"or mtp_step + 1={draft_step + 1}" ) + token_ids[selected_rows, 1:] = block_token_ids[:, bonus_rows:] return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=draft_mem_indexes_cpu, draft_probs=None, ) - @staticmethod - def select_draft_token_ids( - *, - block_token_ids: torch.Tensor, - draft_step: int, - layout: BlockDraftLayout, - ) -> torch.Tensor: - output_start = layout.proposal_output_start - output_end = output_start + int(draft_step) - assert 0 <= output_start <= output_end <= block_token_ids.shape[1] - return block_token_ids[:, output_start:output_end] - - def extend_draft_kv_cache(self, *, main_model_input: ModelInput) -> None: - target_hidden = self.runtime.get_hidden() + def extend_draft_kv_cache(self, main_model_input: ModelInput, target_hidden: torch.Tensor) -> None: draft_model = self.backend.draft_models[0] batch_size = int(target_hidden.shape[0]) assert batch_size == main_model_input.b_req_idx.shape[0] @@ -158,11 +142,9 @@ def extend_draft_kv_cache(self, *, main_model_input: ModelInput) -> None: draft_kv_input.b_prefill_has_output_cpu = [False for _ in range(batch_size)] draft_kv_input.mtp_draft_input_hiddens = target_hidden draft_model.forward(draft_kv_input) - return def build_block_draft_input( self, - *, main_model_input: ModelInput, next_token_ids: torch.Tensor, selected_rows: torch.Tensor, @@ -195,6 +177,7 @@ def build_block_draft_input( draft_input.max_q_seq_len = 1 draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size draft_input.max_cache_len = draft_input.max_kv_seq_len + draft_input.draft_step = block_size - 1 draft_input.b_req_idx = ( main_model_input.b_req_idx.index_select(0, selected_rows).repeat_interleave(block_size).contiguous() ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/speculative/proposers/dspark.py index 1b32a162c0..ca68069835 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dspark.py @@ -3,6 +3,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer @@ -15,50 +16,52 @@ class DSparkProposer(DFlashProposer): confidence logits, so the proposer follows the same token path as DFlash. """ - variant = "dspark" - @torch.no_grad() def propose_next( self, - *, main_model_input: ModelInput, - main_model_output: ModelOutput = None, + main_model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, - verify_result=None, + accept_len: torch.Tensor | None = None, ) -> SpecProposal: - del main_model_output - assert 0 <= draft_step <= self.backend.mtp_step - assert verify_result is not None, "DSpark proposal requires target verify result" + assert 0 <= draft_step <= self.backend.max_draft_step + assert accept_len is not None, "DSpark proposal requires target accept lengths" num_reqs = int(b_req_mtp_start_loc.shape[0]) verify_row_count = next_token_ids.shape[0] draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - layout = self.backend.block_draft_layout - assert block_size == layout.query_block_size - assert verify_result.accept_len.shape[0] == num_reqs + assert block_size == self.backend.max_draft_step, ( + f"DSpark requires --mtp_step={block_size} for this checkpoint, " f"got {self.backend.max_draft_step}" + ) + assert accept_len.shape[0] == num_reqs proposal_token_ids = next_token_ids.new_full( (verify_row_count, draft_step + 1), fill_value=1, ) proposal_token_ids[:, 0] = next_token_ids - schedule_probs = [] if self.enable_dynamic_mtp else None + schedule_probs = [] if self.enable_dynamic_spec else None + + assert main_model_output is not None and main_model_output.spec_hidden is not None + self.extend_draft_kv_cache( + main_model_input=main_model_input, + target_hidden=main_model_output.spec_hidden, + ) if draft_step == 0: return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, - draft_probs=[] if self.enable_dynamic_mtp else None, + draft_probs=[] if self.enable_dynamic_spec else None, schedule_probs=schedule_probs, ) - self.extend_draft_kv_cache(main_model_input=main_model_input) selected_rows = self.select_accepted_tail_rows( b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=verify_result.accept_len, + accept_len=accept_len, ) draft_input, draft_mem_indexes_cpu = self.build_block_draft_input( main_model_input=main_model_input, @@ -67,26 +70,26 @@ def propose_next( num_reqs=num_reqs, ) draft_model_output = draft_model.forward(draft_input) + assert isinstance(draft_model_output, DSparkModelOutput) expected_block_rows = num_reqs * block_size assert draft_model_output.logits.ndim >= 2, "draft logits must have a leading block-row dimension" assert ( draft_model_output.logits.shape[0] == expected_block_rows ), f"draft logits rows must be {expected_block_rows}, got {draft_model_output.logits.shape[0]}" - flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + if draft_model_output.draft_token_ids is None: + flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + else: + flat_token_ids = draft_model_output.draft_token_ids assert ( flat_token_ids.numel() == expected_block_rows ), f"draft token rows must be {expected_block_rows}, got {flat_token_ids.numel()}" block_token_ids = flat_token_ids.reshape(num_reqs, block_size) - proposal_token_ids[selected_rows, 1:] = self.select_draft_token_ids( - block_token_ids=block_token_ids, - draft_step=draft_step, - layout=layout, - ) + proposal_token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] draft_probs = None - if self.enable_dynamic_mtp: - confidence_logits = draft_model_output.mtp_draft_confidence_logits + if self.enable_dynamic_spec: + confidence_logits = draft_model_output.confidence_logits if confidence_logits is None: raise RuntimeError("DSpark dynamic verify requires confidence head logits") assert confidence_logits.ndim == 2, "confidence logits must be [selected_rows, block_size]" @@ -94,13 +97,11 @@ def propose_next( confidence_logits.shape[0] == num_reqs ), f"confidence logits rows must be {num_reqs}, got {confidence_logits.shape[0]}" assert ( - confidence_logits.shape[1] >= layout.proposal_output_start + draft_step - ), f"confidence logits columns must cover the proposal layout, got {confidence_logits.shape[1]}" - output_start = layout.proposal_output_start - output_end = output_start + draft_step + confidence_logits.shape[1] == block_size + ), f"confidence logits must have {block_size} columns, got {confidence_logits.shape[1]}" schedule_probs = self._scatter_step_probs( selected_rows=selected_rows, - probs=confidence_logits[:, output_start:output_end].sigmoid(), + probs=confidence_logits[:, :draft_step].sigmoid(), verify_row_count=verify_row_count, ) @@ -113,7 +114,6 @@ def propose_next( def _scatter_step_probs( self, - *, selected_rows: torch.Tensor, probs: torch.Tensor, verify_row_count: int, diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py index f2be684106..df380c83c5 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py @@ -18,15 +18,18 @@ class Eagle3Proposer(RecurrentEagleMTPProposer): from that corrected state. """ - def __init__(self, runtime) -> None: - super().__init__(runtime) - self._confidence_draft_prune = os.getenv( - # Dynamic Eagle3 already ranks target verify rows by proposal - # confidence. Apply the same budget to deep draft frontiers by - # default; static Eagle3 is explicitly excluded below. - "LIGHTLLM_EAGLE3_CONFIDENCE_DRAFT_PRUNE", - "1", - ).lower() in {"1", "true", "yes", "on"} + def __init__(self, engine) -> None: + super().__init__(engine) + self._confidence_draft_prune = ( + os.getenv( + # Dynamic Eagle3 already ranks target verify rows by proposal + # confidence. Apply the same budget to deep draft frontiers by + # default; static Eagle3 is explicitly excluded below. + "LIGHTLLM_EAGLE3_CONFIDENCE_DRAFT_PRUNE", + "1", + ).lower() + in {"1", "true", "yes", "on"} + ) self._draft_prune_safety_factor = max( 0.0, float(os.getenv("LIGHTLLM_EAGLE3_DRAFT_PRUNE_SAFETY_FACTOR", "1.10")), @@ -38,7 +41,6 @@ def __init__(self, runtime) -> None: def _get_pruned_active_count( self, - *, current_count: int, draft_row_budget: int, next_depth: int, @@ -55,17 +57,16 @@ def _get_pruned_active_count( def propose_next( self, - *, main_model_input: ModelInput, - main_model_output: ModelOutput = None, + main_model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, - verify_result=None, + accept_len: torch.Tensor | None = None, ) -> SpecProposal: - assert 0 <= draft_step <= self.backend.mtp_step - assert verify_result is not None, "Eagle3 proposal requires target verify result" - del main_model_output + assert 0 <= draft_step <= self.backend.max_draft_step + assert accept_len is not None, "Eagle3 proposal requires target accept lengths" + assert main_model_output is not None and main_model_output.spec_hidden is not None verify_row_count = next_token_ids.shape[0] num_reqs = b_req_mtp_start_loc.shape[0] proposal_token_ids = next_token_ids.new_full( @@ -73,20 +74,11 @@ def propose_next( fill_value=1, ) proposal_token_ids[:, 0].copy_(next_token_ids) - collect_dynamic_probs = self.enable_dynamic_mtp and self.runtime.collect_dynamic_probs + collect_dynamic_probs = self.enable_dynamic_spec draft_probs = [] if collect_dynamic_probs else None - if draft_step == 0: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - draft_probs=draft_probs, - draft_forward_rows=0, - ) + target_hidden = main_model_output.spec_hidden - target_hidden = self.runtime.get_hidden() - - accept_len = verify_result.accept_len # Scatter consumes the accepted-tail row for each request; only those # rows need new draft columns after the commit step. selected_rows = self.select_accepted_tail_rows( @@ -94,12 +86,20 @@ def propose_next( accept_len=accept_len, ) draft_model = self.backend.draft_models[0] - draft_model_input = self.make_full_verify_decode_input( + draft_model_input = self.make_verify_extend_input( base_input=main_model_input, input_ids=next_token_ids, draft_hidden=target_hidden, ) draft_model_output = draft_model.forward(draft_model_input) + + if draft_step == 0: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + draft_probs=draft_probs, + ) + draft_logits = draft_model_output.logits.index_select(0, selected_rows) if collect_dynamic_probs: draft_next_token_ids, selected_draft_prob = self.backend._gen_argmax_token_ids_and_prob( @@ -115,16 +115,14 @@ def propose_next( else: draft_next_token_ids = self.backend._gen_argmax_token_ids(ModelOutput(logits=draft_logits)) chain_survival = None - draft_hidden = self.runtime.get_hidden().index_select(0, selected_rows) + assert draft_model_output.spec_hidden is not None + draft_hidden = draft_model_output.spec_hidden.index_select(0, selected_rows) proposal_token_ids[selected_rows, 1] = draft_next_token_ids - draft_forward_rows = verify_row_count - if draft_step == 1: return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, draft_probs=draft_probs, - draft_forward_rows=draft_forward_rows, ) eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * (draft_step - 1)) @@ -186,7 +184,6 @@ def propose_next( max_kv_seq_len=main_model_input.max_kv_seq_len + step, ) draft_output = draft_model.forward(draft_input) - draft_forward_rows += active_count if collect_dynamic_probs: draft_next_token_ids, selected_draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_output) draft_prob = self.scatter_selected_step_probs( @@ -199,12 +196,12 @@ def propose_next( else: draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_output) proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids - draft_hidden = self.runtime.get_hidden() + draft_hidden = draft_output.spec_hidden + assert draft_hidden is not None selected_seq_len = selected_seq_len + 1 return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=eagle_mem_indexes_cpu, draft_probs=draft_probs, - draft_forward_rows=draft_forward_rows, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py index 9efc6003e2..b93e7ca51d 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -14,76 +14,77 @@ class RecurrentEagleMTPProposer(VanillaMTPProposer): def build_initial_draft_state( self, - *, model_input: ModelInput, + model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: draft_model_input = self.prepare_draft_prefill_input( model_input=model_input, next_token_ids=next_token_ids, + mtp_draft_input_hiddens=model_output.spec_hidden, ) self.backend.draft_models[0].forward(draft_model_input) - return None def build_initial_draft_state_overlap( self, - *, model_input0: ModelInput, + model_output0: ModelOutput, next_token_ids0: torch.Tensor, model_input1: ModelInput, + model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: draft_model_input0 = self.prepare_draft_prefill_input( model_input=model_input0, next_token_ids=next_token_ids0, - microbatch_index=0, + mtp_draft_input_hiddens=model_output0.spec_hidden, ) draft_model_input1 = self.prepare_draft_prefill_input( model_input=model_input1, next_token_ids=next_token_ids1, - microbatch_index=1, + mtp_draft_input_hiddens=model_output1.spec_hidden, ) self.backend.draft_models[0].microbatch_overlap_prefill(draft_model_input0, draft_model_input1) - return None def project_draft_decode_hidden(self, draft_hidden: torch.Tensor) -> torch.Tensor: - if draft_hidden is None: - return None + return draft_hidden - draft_model = self.backend.draft_models[0] - pre_infer = getattr(draft_model, "pre_infer", None) - projector = getattr(pre_infer, "project_mtp_draft_hiddens", None) - if projector is None: - return draft_hidden - - # Keep draft decode CUDA graph input shape stable across target and draft hidden sources. - return projector( - draft_hidden, - draft_model.pre_post_weight, - use_custom_tensor_mananger=False, - ) - - def make_full_verify_decode_input( + def make_verify_extend_input( self, - *, base_input: ModelInput, input_ids: torch.Tensor, draft_hidden: torch.Tensor, ) -> ModelInput: new_input = copy.copy(base_input) + batch_size = int(input_ids.shape[0]) + assert batch_size == int(draft_hidden.shape[0]) + assert batch_size == int(base_input.b_seq_len.shape[0]) + new_input.is_prefill = True + new_input.batch_size = batch_size + new_input.total_token_num = batch_size + new_input.prefix_total_token_num = 0 + new_input.max_q_seq_len = 1 + new_input.max_cache_len = max(0, int(base_input.max_cache_len or 0)) new_input.input_ids = input_ids new_input.mtp_draft_input_hiddens = draft_hidden new_input.mem_indexes_cpu = None - new_input.disable_mtp_decode_att = False - # The target may be using the fixed-layout full fast path. Draft - # forwards have their own graph/layout policy and must not inherit - # that target-only flag through the shallow ModelInput copy. - new_input.use_static_mtp_layout = False + # Treat each verify row as a one-token extend request. This writes the + # accepted prefix into draft KV without giving recurrent decode a + # second attention topology or CUDA graph variant. + new_input.b_ready_cache_len = base_input.b_seq_len - 1 + new_input.b_prefill_start_loc = torch.arange( + batch_size, + dtype=torch.int32, + device=input_ids.device, + ) + new_input.b_prefill_has_output_cpu = [False] * batch_size + new_input.b_position_delta = None + new_input.b_is_decode_req = None + new_input.multimodal_params = [{"images": [], "audios": []}] * batch_size return new_input def make_single_step_decode_input( self, - *, base_input: ModelInput, input_ids: torch.Tensor, draft_hidden: torch.Tensor, @@ -107,11 +108,10 @@ def make_single_step_decode_input( new_input.mem_indexes_cpu = None new_input.b_mark_shared_group = b_mark_shared_group new_input.b_shared_seq_len = None - new_input.disable_mtp_decode_att = True - new_input.use_static_mtp_layout = False new_input.max_q_seq_len = 1 new_input.max_kv_seq_len = max_kv_seq_len new_input.total_token_num = new_input.batch_size * max_kv_seq_len + new_input.draft_step = 0 # Recurrent Eagle decode only needs a correctly sized placeholder # list. Nested per-row allocations otherwise sit between graph # replays and extend the draft proposal critical path. @@ -119,6 +119,195 @@ def make_single_step_decode_input( new_input.multimodal_params = [empty_multimodal_params] * new_input.batch_size return new_input + @staticmethod + def _pad_step_mem_indexes( + real_mem_indexes: torch.Tensor, + request_capacity: int, + hold_mem_index: int, + ) -> torch.Tensor: + padded = torch.full( + (request_capacity,), + hold_mem_index, + dtype=real_mem_indexes.dtype, + device=real_mem_indexes.device, + ) + padded[: real_mem_indexes.shape[0]].copy_(real_mem_indexes) + return padded + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor, + draft_step: int, + ) -> SpecProposal: + """Run recurrent Eagle with DP microbatch overlap. + + Target verification remains physically padded to ``B * (K + 1)`` for + DP collectives. Draft state is corrected once over those verify rows; + recurrent draft decode then runs one row per logical request (plus + HOLD rows required to keep the two microbatches shape-compatible). + """ + + assert 0 <= draft_step <= self.backend.max_draft_step + verify_width = self.backend.max_draft_step + 1 + inputs = (main_model_input0, main_model_input1) + outputs = (main_model_output0, main_model_output1) + next_ids = (next_token_ids0, next_token_ids1) + real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) + accept_lens = (accept_len0, accept_len1) + + request_capacities = [] + real_request_nums = [] + selected_rows = [] + extend_inputs = [] + for model_input, model_output, token_ids, real_rows, accept_len in zip( + inputs, outputs, next_ids, real_verify_rows, accept_lens + ): + assert model_output.spec_hidden is not None + assert model_input.batch_size % verify_width == 0 + assert real_rows % verify_width == 0 + assert token_ids.shape[0] == model_input.batch_size + request_capacity = model_input.batch_size // verify_width + real_request_num = real_rows // verify_width + assert accept_len.shape[0] == request_capacity + starts = torch.arange( + 0, + model_input.batch_size, + verify_width, + dtype=torch.int32, + device=token_ids.device, + ) + selected = self.select_accepted_tail_rows( + b_req_mtp_start_loc=starts, + accept_len=accept_len, + ) + request_capacities.append(request_capacity) + real_request_nums.append(real_request_num) + selected_rows.append(selected) + extend_inputs.append( + self.make_verify_extend_input( + base_input=model_input, + input_ids=token_ids, + draft_hidden=model_output.spec_hidden, + ) + ) + + real_row_count = real_verify_rows0 + real_verify_rows1 + proposal_token_ids = next_token_ids0.new_full( + (real_row_count, draft_step + 1), + fill_value=1, + ) + proposal_token_ids[:real_verify_rows0, 0].copy_(next_token_ids0[:real_verify_rows0]) + proposal_token_ids[real_verify_rows0:, 0].copy_(next_token_ids1[:real_verify_rows1]) + + draft_model = self.backend.draft_models[0] + extend_outputs = draft_model.microbatch_overlap_prefill(*extend_inputs) + if draft_step == 0: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + ) + + draft_next_token_ids = [] + draft_hiddens = [] + selected_seq_lens = [] + selected_req_idxs = [] + selected_mtp_idxs = [] + selected_position_deltas = [] + group_marks = [] + proposal_row_offsets = (0, real_verify_rows0) + for index, (model_input, extend_output, selected, real_request_num) in enumerate( + zip(inputs, extend_outputs, selected_rows, real_request_nums) + ): + selected_output = ModelOutput(logits=extend_output.logits.index_select(0, selected)) + step_token_ids = self.backend._gen_argmax_token_ids(selected_output) + draft_next_token_ids.append(step_token_ids) + assert extend_output.spec_hidden is not None + draft_hiddens.append(extend_output.spec_hidden.index_select(0, selected)) + selected_seq_lens.append(model_input.b_seq_len.index_select(0, selected) + 1) + selected_req_idxs.append(model_input.b_req_idx.index_select(0, selected)) + selected_mtp_idxs.append(torch.zeros_like(selected_req_idxs[-1])) + selected_position_deltas.append( + model_input.b_position_delta.index_select(0, selected) + if model_input.b_position_delta is not None + else None + ) + group_marks.append( + torch.ones( + request_capacities[index], + dtype=torch.int32, + device=step_token_ids.device, + ) + ) + if real_request_num > 0: + proposal_rows = selected[:real_request_num].to(torch.long) + proposal_row_offsets[index] + proposal_token_ids[proposal_rows, 1] = step_token_ids[:real_request_num] + + if draft_step == 1: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + ) + + total_real_requests = sum(real_request_nums) + extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(total_real_requests * (draft_step - 1)) + extra_mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) + hold_mem_index = self.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX + + for step in range(1, draft_step): + mem_start = (step - 1) * total_real_requests + step_mem_indexes = extra_mem_indexes[mem_start : mem_start + total_real_requests] + step_inputs = [] + real_mem_start = 0 + for index, model_input in enumerate(inputs): + real_request_num = real_request_nums[index] + real_mem_indexes = step_mem_indexes[real_mem_start : real_mem_start + real_request_num] + real_mem_start += real_request_num + padded_mem_indexes = self._pad_step_mem_indexes( + real_mem_indexes=real_mem_indexes, + request_capacity=request_capacities[index], + hold_mem_index=hold_mem_index, + ) + step_inputs.append( + self.make_single_step_decode_input( + base_input=model_input, + input_ids=draft_next_token_ids[index], + draft_hidden=self.project_draft_decode_hidden(draft_hiddens[index]), + b_req_idx=selected_req_idxs[index], + b_mtp_index=selected_mtp_idxs[index], + b_seq_len=selected_seq_lens[index], + b_position_delta=selected_position_deltas[index], + mem_indexes=padded_mem_indexes, + b_mark_shared_group=group_marks[index], + max_kv_seq_len=model_input.max_kv_seq_len + step, + ) + ) + + step_outputs = draft_model.microbatch_overlap_decode(*step_inputs) + for index, step_output in enumerate(step_outputs): + step_token_ids = self.backend._gen_argmax_token_ids(step_output) + draft_next_token_ids[index] = step_token_ids + draft_hiddens[index] = step_output.spec_hidden + assert draft_hiddens[index] is not None + selected_seq_lens[index] = selected_seq_lens[index] + 1 + real_request_num = real_request_nums[index] + if real_request_num > 0: + proposal_rows = selected_rows[index][:real_request_num].to(torch.long) + proposal_row_offsets[index] + proposal_token_ids[proposal_rows, step + 1] = step_token_ids[:real_request_num] + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, + ) + class EagleMTPProposer(RecurrentEagleMTPProposer): """Recurrent Eagle MTP proposer. @@ -129,91 +318,131 @@ class EagleMTPProposer(RecurrentEagleMTPProposer): def propose_next( self, - *, main_model_input: ModelInput, - main_model_output: ModelOutput = None, + main_model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, - verify_result=None, + accept_len: torch.Tensor | None = None, ) -> SpecProposal: - assert 0 <= draft_step <= self.backend.mtp_step - del verify_result - return self._propose_expanded( + assert 0 <= draft_step <= self.backend.max_draft_step + assert accept_len is not None + return self._propose_recurrent( main_model_input=main_model_input, main_model_output=main_model_output, next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, + accept_len=accept_len, ) - def _propose_expanded( + def _propose_recurrent( self, - *, main_model_input: ModelInput, - main_model_output: ModelOutput = None, + main_model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, + accept_len: torch.Tensor, ) -> SpecProposal: - del main_model_output - num_reqs = b_req_mtp_start_loc.shape[0] + assert main_model_output is not None and main_model_output.spec_hidden is not None + verify_row_count = int(next_token_ids.shape[0]) + num_reqs = int(b_req_mtp_start_loc.shape[0]) + selected_rows = self.select_accepted_tail_rows( + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + ) + proposal_token_ids = next_token_ids.new_full( + (verify_row_count, draft_step + 1), + fill_value=1, + ) + proposal_token_ids[:, 0].copy_(next_token_ids) + draft_probs = [] if self.enable_dynamic_spec else None + + draft_model = self.backend.draft_models[0] + extend_input = self.make_verify_extend_input( + base_input=main_model_input, + input_ids=next_token_ids, + draft_hidden=main_model_output.spec_hidden, + ) + extend_output = draft_model.forward(extend_input) if draft_step == 0: - eagle_mem_indexes_cpu = None - eagle_mem_indexes = None - else: - eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * draft_step) - eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) - - draft_model_input = main_model_input - draft_next_token_ids = next_token_ids - draft_hidden = self.runtime.get_hidden() if draft_step > 0 else None - all_next_token_ids = [next_token_ids] - draft_probs = [] if self.enable_dynamic_mtp else None - - for step in range(draft_step): - draft_hidden = self.project_draft_decode_hidden(draft_hidden) - draft_model_input = self.prepare_draft_decode_input( - model_input=draft_model_input, - next_token_ids=draft_next_token_ids, - mtp_draft_input_hiddens=draft_hidden, + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + draft_probs=draft_probs, ) - draft_model = self.backend.draft_models[0] - draft_model_output = draft_model.forward(draft_model_input) - draft_hidden = self.runtime.get_hidden() - if self.enable_dynamic_mtp: - draft_next_token_ids, draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) - draft_probs.append(draft_prob) - else: - draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) - - draft_model_input.b_seq_len += 1 - draft_model_input.max_kv_seq_len += 1 - eagle_mem_indexes_i = eagle_mem_indexes[step * num_reqs : (step + 1) * num_reqs] - if self.enable_dynamic_mtp: - from lightllm.server.router.model_infer.mode_backend.update_mem_index import ( - update_eagle_mem_indexes_triton, + selected_logits = extend_output.logits.index_select(0, selected_rows) + selected_output = ModelOutput(logits=selected_logits) + if self.enable_dynamic_spec: + draft_next_token_ids, selected_prob = self.backend._gen_argmax_token_ids_and_prob(selected_output) + draft_probs.append( + self.scatter_selected_step_probs( + selected_rows=selected_rows, + selected_probs=selected_prob, + verify_row_count=verify_row_count, ) + ) + else: + draft_next_token_ids = self.backend._gen_argmax_token_ids(selected_output) + proposal_token_ids[selected_rows, 1] = draft_next_token_ids + assert extend_output.spec_hidden is not None + draft_hidden = extend_output.spec_hidden.index_select(0, selected_rows) + + if draft_step == 1: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + draft_probs=draft_probs, + ) - draft_model_input.mem_indexes = update_eagle_mem_indexes_triton( - old_mem_indexes=draft_model_input.mem_indexes, - new_step_mem_indexes=eagle_mem_indexes_i, - b_req_mtp_start_loc=b_req_mtp_start_loc, + eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * (draft_step - 1)) + eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) + selected_seq_len = main_model_input.b_seq_len.index_select(0, selected_rows) + 1 + selected_req_idx = main_model_input.b_req_idx.index_select(0, selected_rows) + selected_mtp_index = torch.zeros_like(selected_req_idx) + selected_position_delta = ( + main_model_input.b_position_delta.index_select(0, selected_rows) + if main_model_input.b_position_delta is not None + else None + ) + one_row_group_marks = torch.ones(num_reqs, dtype=torch.int32, device=next_token_ids.device) + + for step in range(1, draft_step): + mem_start = (step - 1) * num_reqs + draft_input = self.make_single_step_decode_input( + base_input=main_model_input, + input_ids=draft_next_token_ids, + draft_hidden=self.project_draft_decode_hidden(draft_hidden), + b_req_idx=selected_req_idx, + b_mtp_index=selected_mtp_index, + b_seq_len=selected_seq_len, + b_position_delta=selected_position_delta, + mem_indexes=eagle_mem_indexes[mem_start : mem_start + num_reqs], + b_mark_shared_group=one_row_group_marks, + max_kv_seq_len=main_model_input.max_kv_seq_len + step, + ) + draft_output = draft_model.forward(draft_input) + if self.enable_dynamic_spec: + draft_next_token_ids, selected_prob = self.backend._gen_argmax_token_ids_and_prob(draft_output) + draft_probs.append( + self.scatter_selected_step_probs( + selected_rows=selected_rows, + selected_probs=selected_prob, + verify_row_count=verify_row_count, + ) ) else: - draft_model_input.mem_indexes = torch.cat( - [ - draft_model_input.mem_indexes.view(-1, self.backend.mtp_step + 1)[:, 1:], - eagle_mem_indexes_i.view(-1, 1), - ], - dim=1, - ).view(-1) - all_next_token_ids.append(draft_next_token_ids) + draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_output) + proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids + draft_hidden = draft_output.spec_hidden + assert draft_hidden is not None + selected_seq_len = selected_seq_len + 1 return SpecProposal( - token_ids=torch.stack(all_next_token_ids, dim=1), + token_ids=proposal_token_ids, extra_mem_indexes_cpu=eagle_mem_indexes_cpu, draft_probs=draft_probs, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py index 4f5db86cdf..7654960259 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py @@ -9,25 +9,26 @@ class VanillaMTPProposer(BaseSpecProposer): """Chained MTP proposer. - This path uses `mtp_step` independent draft modules. Step i consumes the + This path uses `max_draft_step` independent draft modules. Step i consumes the hidden feature produced by step i - 1 and predicts one candidate token. Target -> draft transfer: - target prefill/decode captures final hidden states with shape [token_num, hidden_size] - - runtime injects them into ModelInput.mtp_draft_input_hiddens before each + - SpecEngine injects them into ModelInput.mtp_draft_input_hiddens before each draft forward """ def build_initial_draft_state( self, - *, model_input: ModelInput, + model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: draft_model_input = model_input draft_next_token_ids = next_token_ids - draft_hidden = self.runtime.get_hidden() + draft_hidden = model_output.spec_hidden + assert draft_hidden is not None for draft_model in self.backend.draft_models: draft_model_input = self.prepare_draft_prefill_input( model_input=draft_model_input, @@ -36,36 +37,36 @@ def build_initial_draft_state( ) draft_model_output = draft_model.forward(draft_model_input) draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) - draft_hidden = self.runtime.get_hidden() - return None + draft_hidden = draft_model_output.spec_hidden + assert draft_hidden is not None def build_initial_draft_state_overlap( self, - *, model_input0: ModelInput, + model_output0: ModelOutput, next_token_ids0: torch.Tensor, model_input1: ModelInput, + model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: draft_model_input0 = model_input0 draft_model_input1 = model_input1 draft_next_token_ids0 = next_token_ids0 draft_next_token_ids1 = next_token_ids1 - draft_hidden0 = self.runtime.get_hidden(0) - draft_hidden1 = self.runtime.get_hidden(1) + draft_hidden0 = model_output0.spec_hidden + draft_hidden1 = model_output1.spec_hidden + assert draft_hidden0 is not None and draft_hidden1 is not None for draft_model in self.backend.draft_models: draft_model_input0 = self.prepare_draft_prefill_input( model_input=draft_model_input0, next_token_ids=draft_next_token_ids0, mtp_draft_input_hiddens=draft_hidden0, - microbatch_index=0, ) draft_model_input1 = self.prepare_draft_prefill_input( model_input=draft_model_input1, next_token_ids=draft_next_token_ids1, mtp_draft_input_hiddens=draft_hidden1, - microbatch_index=1, ) draft_model_output0, draft_model_output1 = draft_model.microbatch_overlap_prefill( draft_model_input0, @@ -73,28 +74,25 @@ def build_initial_draft_state_overlap( ) draft_next_token_ids0 = self.backend._gen_argmax_token_ids(draft_model_output0) draft_next_token_ids1 = self.backend._gen_argmax_token_ids(draft_model_output1) - draft_hidden0 = self.runtime.get_hidden(0) - draft_hidden1 = self.runtime.get_hidden(1) - return None + draft_hidden0 = draft_model_output0.spec_hidden + draft_hidden1 = draft_model_output1.spec_hidden + assert draft_hidden0 is not None and draft_hidden1 is not None def propose_next( self, - *, main_model_input: ModelInput, - main_model_output: ModelOutput = None, + main_model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, - verify_result=None, + accept_len: torch.Tensor | None = None, ) -> SpecProposal: - del main_model_output - del verify_result - assert 0 <= draft_step <= self.backend.mtp_step + assert 0 <= draft_step <= self.backend.max_draft_step draft_model_input = main_model_input draft_next_token_ids = next_token_ids - draft_hidden = self.runtime.get_hidden() if draft_step > 0 else None + draft_hidden = main_model_output.spec_hidden if draft_step > 0 else None all_next_token_ids = [next_token_ids] - draft_probs = [] if self.enable_dynamic_mtp else None + draft_probs = [] if self.enable_dynamic_spec else None for step in range(draft_step): draft_model = self.backend.draft_models[step] @@ -104,8 +102,9 @@ def propose_next( mtp_draft_input_hiddens=draft_hidden, ) draft_model_output = draft_model.forward(draft_model_input) - draft_hidden = self.runtime.get_hidden() - if self.enable_dynamic_mtp: + draft_hidden = draft_model_output.spec_hidden + assert draft_hidden is not None + if self.enable_dynamic_spec: draft_next_token_ids, draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) draft_probs.append(draft_prob) else: diff --git a/lightllm/server/router/model_infer/speculative/runner.py b/lightllm/server/router/model_infer/speculative/runner.py index d3cc2ea882..a34e9d9b35 100644 --- a/lightllm/server/router/model_infer/speculative/runner.py +++ b/lightllm/server/router/model_infer/speculative/runner.py @@ -8,14 +8,14 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.mtp_utils import ( gen_b_req_mtp_start_loc, - linear_att_mtp_state_index_update, + linear_att_spec_state_index_update, ) from lightllm.server.router.model_infer.infer_batch import g_infer_context from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager if TYPE_CHECKING: from lightllm.server.router.model_infer.speculative.planner import SpecDecodePlan - from lightllm.server.router.model_infer.speculative.runtime import SpecRuntime + from lightllm.server.router.model_infer.speculative.engine import SpecEngine @dataclass @@ -23,15 +23,15 @@ class SpecDecodeForwardState: model_input: ModelInput original_run_reqs: List plan: "SpecDecodePlan" - selected_run_reqs_cpu: Optional[torch.Tensor] + selected_row_mask_cpu: Optional[torch.Tensor] accepted_index_cpu: torch.Tensor - mtp_accept_len_cpu: torch.Tensor + spec_accept_len_cpu: torch.Tensor next_token_ids_cpu: torch.Tensor next_token_logprobs_cpu: torch.Tensor next_token_ranks_cpu: torch.Tensor verify_event: torch.cuda.Event sync_event: torch.cuda.Event - additional_mem_indexes_cpu: Optional[torch.Tensor] + extra_mem_indexes_cpu: Optional[torch.Tensor] schedule_probs_cpu: Optional[torch.Tensor] @@ -40,24 +40,23 @@ class SpecDecodePostState: next_token_ids: torch.Tensor next_token_logprobs: torch.Tensor next_token_ranks: torch.Tensor - mtp_accept_len_cpu: torch.Tensor + spec_accept_len_cpu: torch.Tensor need_free_mem_indexes: torch.Tensor class SpecDecodeRunner: - def __init__(self, runtime: "SpecRuntime") -> None: - self.runtime = runtime - self.backend = runtime.backend + def __init__(self, engine: "SpecEngine") -> None: + self.engine = engine + self.backend = engine.backend def run_speculative_forward( self, - *, model_input: ModelInput, model_output: ModelOutput, run_reqs: List, req_num: int, plan: "SpecDecodePlan", - selected_run_reqs_cpu: Optional[torch.Tensor], + selected_row_mask_cpu: Optional[torch.Tensor], next_token_ids: torch.Tensor, next_token_logprobs: torch.Tensor, next_token_ranks: torch.Tensor, @@ -66,22 +65,22 @@ def run_speculative_forward( Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ], ) -> SpecDecodeForwardState: - runtime = self.runtime + engine = self.engine b_req_mtp_start_loc = gen_b_req_mtp_start_loc(model_input.b_mtp_index, num_reqs=req_num) - verify_result = runtime.verify_target_tokens( + verify_result = engine.verify_target_tokens( new_next_token_ids=next_token_ids, b_req_idx=model_input.b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, ) accepted_index = verify_result.accepted_index if self.backend.is_linear_att_mixed_model: - linear_att_mtp_state_index_update( + linear_att_spec_state_index_update( req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx, b_mtp_index=model_input.b_mtp_index, accepted_index=accepted_index, - max_mtp_step=self.backend.mtp_step + 1, + verify_width=self.backend.max_draft_step + 1, ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", @@ -91,22 +90,20 @@ def run_speculative_forward( verify_event = torch.cuda.Event(enable_timing=True) verify_event.record() - runtime.configure_dynamic_prob_collection(plan=plan) - proposal = runtime.propose_next( + proposal = engine.propose_next( main_model_input=model_input, main_model_output=model_output, next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=plan.draft_step, - verify_result=verify_result, + accept_len=verify_result.accept_len, ) - all_next_token_ids = runtime.pad_all_next_token_ids( + all_next_token_ids = engine.pad_all_next_token_ids( token_ids=proposal.token_ids, draft_step=plan.draft_step, ) - all_next_token_probs = runtime.build_all_next_token_probs( - next_token_logprobs=next_token_logprobs, + all_next_token_probs = engine.build_all_next_token_probs( proposal=proposal, draft_step=plan.draft_step, ) @@ -115,15 +112,15 @@ def run_speculative_forward( key="mtp_schedule_probs", gpu_tensor=all_next_token_probs, ) - if all_next_token_probs is not None and runtime.needs_schedule_probs_cpu() + if all_next_token_probs is not None and engine.needs_schedule_probs_cpu() else None ) - runtime.scatter_next_tokens( + engine.scatter_next_tokens( b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=all_next_token_ids, b_req_idx=model_input.b_req_idx, - mtp_accept_len=verify_result.accept_len, + spec_accept_len=verify_result.accept_len, all_next_token_probs=all_next_token_probs, ) @@ -139,8 +136,8 @@ def run_speculative_forward( mask=accepted_index == 1, ) - mtp_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="mtp_accept_len", + spec_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="spec_accept_len", gpu_tensor=verify_result.accept_len, ) @@ -151,63 +148,62 @@ def run_speculative_forward( model_input=model_input, original_run_reqs=run_reqs, plan=plan, - selected_run_reqs_cpu=selected_run_reqs_cpu, + selected_row_mask_cpu=selected_row_mask_cpu, accepted_index_cpu=accepted_index_cpu, - mtp_accept_len_cpu=mtp_accept_len_cpu, + spec_accept_len_cpu=spec_accept_len_cpu, next_token_ids_cpu=next_token_ids_cpu, next_token_logprobs_cpu=next_token_logprobs_cpu, next_token_ranks_cpu=next_token_ranks_cpu, verify_event=verify_event, sync_event=sync_event, - additional_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, schedule_probs_cpu=schedule_probs_cpu, ) - def resolve_pre_post_reqs(self, *, state: SpecDecodeForwardState, decode_reqs: List): + def resolve_pre_post_reqs(self, state: SpecDecodeForwardState, decode_reqs: List): if state.plan.skip_verify_sync: - assert self.runtime.enable_dynamic_mtp, "skip_verify_sync should only be True when dynamic MTP is enabled" + assert self.engine.enable_dynamic_spec, "skip_verify_sync requires dynamic speculative scheduling" return decode_reqs, decode_reqs state.verify_event.synchronize() - return self.runtime.build_decode_req_lists( + return self.engine.build_decode_req_lists( original_run_reqs=state.original_run_reqs, - selected_run_reqs_cpu=state.selected_run_reqs_cpu, + selected_row_mask_cpu=state.selected_row_mask_cpu, accepted_index_cpu=state.accepted_index_cpu, ) - def finish_post(self, *, state: SpecDecodeForwardState, req_num: int, run_reqs: List) -> SpecDecodePostState: + def finish_post(self, state: SpecDecodeForwardState, req_num: int, run_reqs: List) -> SpecDecodePostState: state.sync_event.synchronize() - runtime = self.runtime - if runtime.enable_dynamic_mtp: - runtime.update_dynamic_accept_stats( + engine = self.engine + if engine.enable_dynamic_spec: + engine.update_dynamic_accept_stats( req_num=req_num, run_reqs=run_reqs, accepted_index_cpu=state.accepted_index_cpu, - mtp_accept_len_cpu=state.mtp_accept_len_cpu, + spec_accept_len_cpu=state.spec_accept_len_cpu, dynamic_batch_size=state.plan.dynamic_batch_size, - verify_step=state.plan.pre_draft_step, - selection_mode=state.plan.selection_mode, + pre_draft_step=state.plan.pre_draft_step, ) if state.schedule_probs_cpu is not None: - runtime.update_dynamic_schedule_stats( + engine.update_dynamic_schedule_stats( req_num=req_num, schedule_probs_cpu=state.schedule_probs_cpu, ) - need_free_mem_indexes = runtime.build_decode_free_mem_indexes_cpu( + need_free_mem_indexes = engine.build_decode_free_mem_indexes_cpu( model_input=state.model_input, - selected_run_reqs_cpu=state.selected_run_reqs_cpu, + selected_row_mask_cpu=state.selected_row_mask_cpu, accepted_index_cpu=state.accepted_index_cpu, ) - if state.additional_mem_indexes_cpu is not None: - need_free_mem_indexes = torch.cat([need_free_mem_indexes, state.additional_mem_indexes_cpu], dim=0) + if state.extra_mem_indexes_cpu is not None: + need_free_mem_indexes = torch.cat([need_free_mem_indexes, state.extra_mem_indexes_cpu], dim=0) select_mask = state.accepted_index_cpu.to(dtype=torch.bool) return SpecDecodePostState( next_token_ids=state.next_token_ids_cpu[select_mask], next_token_logprobs=state.next_token_logprobs_cpu[select_mask], next_token_ranks=state.next_token_ranks_cpu[select_mask], - mtp_accept_len_cpu=state.mtp_accept_len_cpu, + spec_accept_len_cpu=state.spec_accept_len_cpu, need_free_mem_indexes=need_free_mem_indexes, ) diff --git a/lightllm/server/router/model_infer/speculative/runtime.py b/lightllm/server/router/model_infer/speculative/runtime.py deleted file mode 100644 index 91cf7c51cd..0000000000 --- a/lightllm/server/router/model_infer/speculative/runtime.py +++ /dev/null @@ -1,775 +0,0 @@ -from __future__ import annotations - -import json -import os -from typing import Callable, List, Optional, Tuple - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.speculative.config import SpeculativeConfig, normalize_speculative_draft_config -from lightllm.server.router.model_infer.speculative.planner import FixedMTPPlanner, SpecDecodePlan -from lightllm.server.router.model_infer.speculative.proposers import build_spec_proposer -from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.speculative.runner import ( - SpecDecodeForwardState, - SpecDecodePostState, - SpecDecodeRunner, -) -from lightllm.server.router.model_infer.speculative.state import SpecForwardContext, SpecHiddenStore -from lightllm.server.router.model_infer.speculative.verifier import SpecVerifier, SpecVerifyResult - - -class SpecRuntime: - """Facade between LightLLM backend code and speculative algorithms. - - The runtime keeps speculative decoding out of BaseModel and - chunked_prefill: - - BaseModel only calls hidden-capture methods. - - proposer implementations own draft-model state/proposal generation. - - SpecVerifier owns service-specific verify/scatter kernels. - - The main target->draft data path is: - 1. target model forward captures hidden features into SpecHiddenStore - 2. runtime injects those features into ModelInput.mtp_draft_input_hiddens - 3. proposer forwards the draft model and returns SpecProposal.token_ids - 4. verifier checks target acceptance and scatters candidates for the next - iteration - """ - - def __init__(self, backend) -> None: - self.backend = backend - self._target_layer_ids: Optional[List[int]] = None - self.hidden_store = SpecHiddenStore(self) - self.verifier = SpecVerifier(backend) - self.proposer = build_spec_proposer(self) - self.decode_runner = SpecDecodeRunner(self) - self.planner = self._build_decode_planner() - self._collect_dynamic_probs = self.enable_dynamic_mtp - self._dynamic_accept_stats_calls = 0 - self._full_fast_stats_count = 0 - self._full_fast_stats_interval = max( - 1, - int(os.getenv("LIGHTLLM_EAGLE3_FULL_FAST_STATS_INTERVAL", "8")), - ) - self._full_fast_accept_floor = float( - os.getenv("LIGHTLLM_EAGLE3_FULL_FAST_ACCEPT_FLOOR", "0.75") - ) - - @property - def spec_config(self) -> SpeculativeConfig: - return self.backend.spec_config - - @property - def enable_dynamic_mtp(self) -> bool: - return self.spec_config.dynamic_verify - - @property - def collect_dynamic_probs(self) -> bool: - return self._collect_dynamic_probs - - def configure_dynamic_prob_collection(self, *, plan: SpecDecodePlan) -> None: - # Eagle3's selected-token probabilities are only used to rank rows for - # confidence compaction. A profitable full-width plan neither ranks - # nor compacts rows, so use Static's cheaper argmax-only proposer. An - # ``observe`` full-width plan still records probabilities to prepare a - # possible transition back to confidence scheduling. - self._collect_dynamic_probs = self.enable_dynamic_mtp and not ( - self.planner.planner_mode == "eagle3" and plan.selection_mode == "full" - ) - - @property - def needs_intermediate_target_hidden(self) -> bool: - return self.spec_config.needs_target_layer_hidden - - def is_draft_model(self, model) -> bool: - return any(model is draft_model for draft_model in self.backend.draft_models) - - def is_block_draft_model(self, model) -> bool: - return bool(getattr(self.spec_config, "uses_block_draft_model", False)) and self.is_draft_model(model) - - def is_draft_forward(self, infer_state, model=None) -> bool: - return ( - infer_state.mtp_draft_input_hiddens is not None - or getattr(infer_state, "is_draft_model", False) - or (model is not None and self.is_draft_model(model)) - ) - - def get_capture_layer_ids(self, model, infer_state) -> List[int]: - if self.is_draft_forward(infer_state, model=model): - if self.spec_config.uses_chained_draft_models or self.spec_config.uses_recurrent_draft_model: - return [model.config["n_layer"] - 1] - return [] - if self.needs_intermediate_target_hidden: - return self._get_target_layer_ids(model) - return [] - - def create_forward_context(self, model, infer_state) -> SpecForwardContext: - return SpecForwardContext(runtime=self, model=model, infer_state=infer_state) - - def capture_hidden( - self, - *, - infer_state, - hidden: torch.Tensor, - final_hidden: torch.Tensor, - ) -> torch.Tensor: - return self.hidden_store.capture_hidden( - infer_state=infer_state, - hidden=hidden, - final_hidden=final_hidden, - ) - - def get_hidden(self, microbatch_index: int = 0) -> torch.Tensor: - return self.hidden_store.get_hidden(microbatch_index) - - def unpad_hidden(self, *, token_num: int, microbatch_index: int = 0) -> None: - self.hidden_store.unpad_hidden(token_num=token_num, microbatch_index=microbatch_index) - return - - def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: - """Allocate speculative draft-owned temporary KV slots.""" - - token_count = int(token_count) - assert token_count >= 0 - if token_count == 0: - return torch.empty((0,), dtype=torch.int32, device="cpu") - - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) - return g_infer_context.req_manager.mem_manager.alloc(token_count) - - def build_padded_next_token_ids( - self, - *, - token_ids: Optional[torch.Tensor], - batch_size: int, - copy_len: int = None, - source_start: int = 0, - device=None, - ) -> torch.Tensor: - """Build a padded draft-token input buffer for padded DP batches.""" - - batch_size = int(batch_size) - source_start = int(source_start) - assert batch_size >= 0 - assert source_start >= 0 - if token_ids is None: - assert copy_len is None or int(copy_len) == 0 - assert device is not None - copy_len = 0 - else: - copy_len = token_ids.shape[0] - source_start if copy_len is None else int(copy_len) - assert copy_len >= 0 - assert source_start + copy_len <= token_ids.shape[0] - if device is None: - device = token_ids.device - assert copy_len <= batch_size - - padded_token_ids = torch.zeros((batch_size,), dtype=torch.int64, device=device) - if copy_len > 0: - padded_token_ids[:copy_len].copy_( - token_ids[source_start : source_start + copy_len], - non_blocking=True, - ) - return padded_token_ids - - def build_padded_eagle_step_mem_indexes( - self, - *, - eagle_mem_indexes: torch.Tensor, - step: int, - real_req_num: int, - padded_req_num: int, - ) -> torch.Tensor: - """Build one padded Eagle scratch-index column for DP decode.""" - - step = int(step) - real_req_num = int(real_req_num) - padded_req_num = int(padded_req_num) - assert step >= 0 - assert real_req_num >= 0 - assert padded_req_num >= 0 - - start = step * real_req_num - end = start + real_req_num - assert end <= eagle_mem_indexes.shape[0] - step_mem_indexes = eagle_mem_indexes[start:end] - if padded_req_num == 0: - return step_mem_indexes - - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - hold_token_memindex = g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX - hold_mem_indexes = torch.full( - (padded_req_num,), - int(hold_token_memindex), - dtype=eagle_mem_indexes.dtype, - device=eagle_mem_indexes.device, - ) - return torch.cat([step_mem_indexes, hold_mem_indexes], dim=0) - - def append_padded_eagle_step_mem_indexes( - self, - *, - model_input: ModelInput, - eagle_mem_indexes: torch.Tensor, - step: int, - real_req_num: int, - padded_req_num: int, - mtp_step: int, - ) -> torch.Tensor: - """Roll one padded Eagle scratch-index column into ModelInput.""" - - mtp_step = int(mtp_step) - assert mtp_step >= 0 - step_mem_indexes = self.build_padded_eagle_step_mem_indexes( - eagle_mem_indexes=eagle_mem_indexes, - step=step, - real_req_num=real_req_num, - padded_req_num=padded_req_num, - ) - grouped_mem_indexes = model_input.mem_indexes.view(-1, mtp_step + 1) - assert grouped_mem_indexes.shape[0] == step_mem_indexes.shape[0] - model_input.mem_indexes = torch.cat( - [grouped_mem_indexes[:, 1:], step_mem_indexes.view(-1, 1)], - dim=1, - ).view(-1) - return model_input.mem_indexes - - def prepare_draft_prefill_input( - self, - *, - model_input: ModelInput, - next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: Optional[torch.Tensor] = None, - microbatch_index: int = 0, - ) -> ModelInput: - """Build draft prefill input from target prefill input. - - `next_token_ids`: [run_req_num] - `mtp_draft_input_hiddens`: captured target feature. The first - dimension matches the target prefill token layout after padding/unpad - handling; the second dimension is either hidden_size or - hidden_size * len(target_layer_ids). - """ - - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - return prepare_mtp_prefill_inputs( - model_input=model_input, - b_next_token_ids=next_token_ids, - mtp_draft_input_hiddens=( - self.get_hidden(microbatch_index) if mtp_draft_input_hiddens is None else mtp_draft_input_hiddens - ), - ) - - def prepare_draft_decode_input( - self, - *, - model_input: ModelInput, - next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: Optional[torch.Tensor] = None, - microbatch_index: int = 0, - ) -> ModelInput: - """Mutate a decode ModelInput for one draft forward. - - `next_token_ids`: [verify_batch] - `mtp_draft_input_hiddens`: [verify_batch, hidden_dim_for_draft] - """ - - model_input.input_ids = next_token_ids - if mtp_draft_input_hiddens is None: - mtp_draft_input_hiddens = self.get_hidden(microbatch_index) - model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens - return model_input - - def graph_cache_key(self, model_context, model=None): - model = getattr(model_context, "model", None) or model - infer_state = getattr(model_context, "infer_state", model_context) - role = "draft" if self.is_draft_forward(infer_state, model=model) else "main" - disable_mtp_decode_att = bool(getattr(infer_state, "disable_mtp_decode_att", False)) - use_static_mtp_layout = bool( - role == "main" - and not disable_mtp_decode_att - and getattr(infer_state, "use_static_mtp_layout", False) - ) - return ("spec", self.spec_config.mode, role, disable_mtp_decode_att, use_static_mtp_layout) - - def get_decode_graph_mtp_step(self, model) -> int: - return self.spec_config.get_decode_graph_mtp_step( - model_config=model.config, - is_draft_model=any(model is draft_model for draft_model in self.backend.draft_models), - ) - - def get_decode_graph_warmup_mtp_step(self, model) -> int: - return self.spec_config.get_decode_graph_warmup_mtp_step( - model_config=model.config, - is_draft_model=any(model is draft_model for draft_model in self.backend.draft_models), - ) - - def export_graph_capture(self): - return self.hidden_store.export_graph_capture() - - def restore_graph_capture(self, captured_hiddens) -> None: - self.hidden_store.restore_graph_capture(captured_hiddens) - return - - def build_initial_draft_state( - self, - *, - model_input: ModelInput, - next_token_ids: torch.Tensor, - ) -> None: - self.proposer.build_initial_draft_state(model_input=model_input, next_token_ids=next_token_ids) - return - - def build_initial_draft_state_overlap( - self, - *, - model_input0: ModelInput, - next_token_ids0: torch.Tensor, - model_input1: ModelInput, - next_token_ids1: torch.Tensor, - ) -> None: - self.proposer.build_initial_draft_state_overlap( - model_input0=model_input0, - next_token_ids0=next_token_ids0, - model_input1=model_input1, - next_token_ids1=next_token_ids1, - ) - return - - def plan_decode(self, *, model_input: ModelInput, req_num: int) -> SpecDecodePlan: - """Return the static or dynamic MTP plan for one decode iteration.""" - - return self.planner.plan(req_num=req_num, original_batch_size=model_input.batch_size) - - def run_decode_speculative_forward( - self, - *, - model_input: ModelInput, - model_output: ModelOutput, - run_reqs: List, - req_num: int, - plan: SpecDecodePlan, - selected_run_reqs_cpu: Optional[torch.Tensor], - next_token_ids: torch.Tensor, - next_token_logprobs: torch.Tensor, - next_token_ranks: torch.Tensor, - copy_next_token_infos: Callable[ - [torch.Tensor, torch.Tensor, torch.Tensor], - Tuple[torch.Tensor, torch.Tensor, torch.Tensor], - ], - ) -> SpecDecodeForwardState: - return self.decode_runner.run_speculative_forward( - model_input=model_input, - model_output=model_output, - run_reqs=run_reqs, - req_num=req_num, - plan=plan, - selected_run_reqs_cpu=selected_run_reqs_cpu, - next_token_ids=next_token_ids, - next_token_logprobs=next_token_logprobs, - next_token_ranks=next_token_ranks, - copy_next_token_infos=copy_next_token_infos, - ) - - def resolve_decode_pre_post_reqs(self, *, state: SpecDecodeForwardState, decode_reqs: List): - return self.decode_runner.resolve_pre_post_reqs(state=state, decode_reqs=decode_reqs) - - def finish_decode_post( - self, - *, - state: SpecDecodeForwardState, - req_num: int, - run_reqs: List, - ) -> SpecDecodePostState: - return self.decode_runner.finish_post(state=state, req_num=req_num, run_reqs=run_reqs) - - def prepare_decode_model_input( - self, - *, - model_input: ModelInput, - req_num: int, - plan: SpecDecodePlan, - ): - """Apply dynamic MTP row compaction when the planner selects it.""" - - if not plan.is_dynamic: - return model_input, None - - model_input.use_static_mtp_layout = False - if plan.selection_mode in {"full", "observe"}: - assert plan.dynamic_batch_size == model_input.batch_size - assert plan.pre_draft_step == self.backend.mtp_step - model_input.use_static_mtp_layout = True - return model_input, None - - self._clear_stale_dynamic_token_probs(pre_draft_step=plan.pre_draft_step) - - from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_mtp_model_input - - selection_verify_step = plan.pre_draft_step - if plan.selection_mode == "prefix": - assert plan.dynamic_batch_size % req_num == 0 - selection_verify_step = plan.dynamic_batch_size // req_num - 1 - - model_input, selected_run_reqs = prepare_dynamic_mtp_model_input( - model_input=model_input, - req_num=req_num, - dynamic_batch_size=plan.dynamic_batch_size, - req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - req_to_next_token_probs=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs, - verify_step=selection_verify_step, - use_prefix_selection=plan.selection_mode == "prefix", - ) - return model_input, selected_run_reqs - - def async_copy_selected_run_reqs(self, selected_run_reqs: Optional[torch.Tensor]): - if selected_run_reqs is None: - return None - from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager - - return g_pin_mem_manager.async_copy_from_gpu_tensor( - key="selected_run_reqs", - gpu_tensor=selected_run_reqs, - ) - - def build_decode_req_lists( - self, - *, - original_run_reqs, - selected_run_reqs_cpu: Optional[torch.Tensor], - accepted_index_cpu: torch.Tensor, - ): - """Build post-handle request lists after optional dynamic MTP compaction.""" - - if self.enable_dynamic_mtp and selected_run_reqs_cpu is not None: - assert selected_run_reqs_cpu is not None - selected_run_reqs_cpu_numpy = selected_run_reqs_cpu.numpy() - run_reqs = [ - original_run_reqs[i] for i in range(len(original_run_reqs)) if selected_run_reqs_cpu_numpy[i] == 1 - ] - else: - run_reqs = original_run_reqs - - accepted_index_cpu_numpy = accepted_index_cpu.numpy() - verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu_numpy[i] == 1] - return run_reqs, verify_ok_reqs - - def build_decode_free_mem_indexes_cpu( - self, - *, - model_input: ModelInput, - selected_run_reqs_cpu: Optional[torch.Tensor], - accepted_index_cpu: torch.Tensor, - ) -> torch.Tensor: - mem_indexes_cpu = model_input.mem_indexes_cpu - if not self.enable_dynamic_mtp or selected_run_reqs_cpu is None: - return mem_indexes_cpu[accepted_index_cpu == 0] - - assert selected_run_reqs_cpu is not None - selected_mask = selected_run_reqs_cpu.to(dtype=torch.bool) - accepted_mask = accepted_index_cpu.to(dtype=torch.bool) - selected_mem_indexes_cpu = mem_indexes_cpu[selected_mask] - assert selected_mem_indexes_cpu.shape[0] == accepted_mask.shape[0] - - unselected_mem_indexes_cpu = mem_indexes_cpu[~selected_mask] - rejected_selected_mem_indexes_cpu = selected_mem_indexes_cpu[~accepted_mask] - if len(unselected_mem_indexes_cpu) == 0: - return rejected_selected_mem_indexes_cpu - if len(rejected_selected_mem_indexes_cpu) == 0: - return unselected_mem_indexes_cpu - return torch.cat([unselected_mem_indexes_cpu, rejected_selected_mem_indexes_cpu], dim=0) - - def update_dynamic_accept_stats( - self, - *, - req_num: int, - run_reqs, - accepted_index_cpu: torch.Tensor, - mtp_accept_len_cpu: torch.Tensor, - dynamic_batch_size: Optional[int], - verify_step: Optional[int] = None, - selection_mode: str = "confidence", - ) -> None: - if not self.enable_dynamic_mtp: - return - - assert dynamic_batch_size is not None - assert len(run_reqs) == accepted_index_cpu.shape[0] - assert mtp_accept_len_cpu.shape[0] == req_num - accept_lengths = mtp_accept_len_cpu.numpy() - accept_count = int(accept_lengths.sum()) - total_count = int(dynamic_batch_size) - is_full_verify = dynamic_batch_size == req_num * (self.backend.mtp_step + 1) - - # The first target decode has no preceding draft proposal, so every - # request structurally accepts only its base token. Treating that - # cold-start iteration as a K-wide acceptance sample makes a highly - # predictable workload look maximally hard and can collapse the - # controller before the first real proposal is verified. - self._dynamic_accept_stats_calls += 1 - if self._dynamic_accept_stats_calls == 1: - return - - # In the profitable full-width state, the per-iteration trace already - # records exact acceptance. Planner EMAs only need periodic refreshes - # while acceptance remains high. A workload drop bypasses sampling - # immediately, so contraction still reacts on the first bad batch. - if self.planner.planner_mode == "eagle3" and selection_mode == "full": - self._full_fast_stats_count += 1 - current_accept_ratio = accept_count / max(1, total_count) - if ( - current_accept_ratio >= self._full_fast_accept_floor - and self._full_fast_stats_count % self._full_fast_stats_interval != 0 - ): - return - - self.planner.update_req_num_to_dynamic_batch_size_to_accept_ratio( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - accept_ratio=accept_count / total_count, - **({"verify_step": verify_step} if self.planner.planner_mode == "eagle3" else {}), - ) - - update_full_verify_tokens_per_req = getattr( - self.planner, - "update_full_verify_tokens_per_req", - None, - ) - if update_full_verify_tokens_per_req is not None and is_full_verify: - update_full_verify_tokens_per_req( - accept_count / req_num, - req_num=req_num, - ) - - update_observed_iteration_stats = getattr( - self.planner, - "update_observed_iteration_stats", - None, - ) - if update_observed_iteration_stats is not None: - update_observed_iteration_stats( - tokens_per_req=accept_count / req_num, - verify_rows_per_req=dynamic_batch_size / req_num, - is_full_verify=is_full_verify, - req_num=req_num, - ) - - # Eagle3 uses these values as an unbiased survival curve. Updating it - # from confidence-selected dynamic rows would bias every depth upward; - # full-width warmup/probe iterations are the valid samples. - update_verified_batch_prefix_stats = getattr( - self.planner, - "update_verified_batch_prefix_stats", - None, - ) - if is_full_verify and update_verified_batch_prefix_stats is not None: - update_verified_batch_prefix_stats( - verify_and_accept_lengths=[ - (self.backend.mtp_step + 1, int(accept_len)) for accept_len in accept_lengths - ], - ) - elif self.planner.planner_mode != "eagle3": - verify_rows_per_req = max(1, int(round(dynamic_batch_size / req_num))) - for accept_len in accept_lengths: - self.planner.update_verified_prefix_stats( - verify_len=verify_rows_per_req, - accept_len=int(accept_len), - ) - return - - def needs_schedule_probs_cpu(self) -> bool: - """Whether this planner consumes proposal confidence on the CPU.""" - - planner_needs_schedule_probs_cpu = getattr(self.planner, "needs_schedule_probs_cpu", None) - if planner_needs_schedule_probs_cpu is not None: - return bool(planner_needs_schedule_probs_cpu()) - return callable(getattr(self.planner, "update_predicted_schedule_probs", None)) - - def update_dynamic_schedule_stats( - self, - *, - req_num: int, - schedule_probs_cpu: Optional[torch.Tensor], - ) -> None: - if not self.enable_dynamic_mtp or schedule_probs_cpu is None: - return - - update_predicted_schedule_probs = getattr(self.planner, "update_predicted_schedule_probs", None) - if update_predicted_schedule_probs is None: - return - - update_predicted_schedule_probs( - schedule_probs=schedule_probs_cpu, - req_num=req_num, - ) - return - - def propose_next( - self, - *, - main_model_input: ModelInput, - main_model_output: Optional[ModelOutput] = None, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - verify_result: Optional[SpecVerifyResult] = None, - ) -> SpecProposal: - return self.proposer.propose_next( - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=draft_step, - verify_result=verify_result, - ) - - def verify_target_tokens( - self, - *, - new_next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - ) -> SpecVerifyResult: - return self.verifier.verify_target_tokens( - new_next_token_ids=new_next_token_ids, - b_req_idx=b_req_idx, - b_req_mtp_start_loc=b_req_mtp_start_loc, - ) - - def build_all_next_token_probs( - self, - *, - next_token_logprobs: torch.Tensor, - proposal: SpecProposal, - draft_step: int, - ) -> Optional[torch.Tensor]: - """Build selected-token probability matrix for dynamic MTP scatter. - - Output shape is [verify_batch, mtp_step + 1]. Column 0 is the target - token probability, fixed to 1 because the target sample is always the - base accepted position. Draft columns store selected-token - probabilities from each proposer step. - """ - - if not self.enable_dynamic_mtp or not self.collect_dynamic_probs: - return None - - schedule_probs = proposal.schedule_probs if proposal.schedule_probs is not None else proposal.draft_probs - assert schedule_probs is not None - - all_next_token_probs = torch.zeros( - size=(next_token_logprobs.shape[0], self.backend.mtp_step + 1), - dtype=torch.float32, - device=next_token_logprobs.device, - ) - all_next_token_probs[:, 0] = 1.0 - - if isinstance(schedule_probs, torch.Tensor): - assert schedule_probs.shape == (next_token_logprobs.shape[0], draft_step) - if draft_step > 0: - all_next_token_probs[:, 1 : draft_step + 1] = schedule_probs - return all_next_token_probs - - assert len(schedule_probs) == draft_step - for step_idx, step_probs in enumerate(schedule_probs): - all_next_token_probs[:, step_idx + 1] = step_probs - return all_next_token_probs - - def pad_all_next_token_ids(self, *, token_ids: torch.Tensor, draft_step: int) -> torch.Tensor: - """Pad dynamic proposal ids to the static MTP width before scatter.""" - - if not self.enable_dynamic_mtp or draft_step >= self.backend.mtp_step: - return token_ids - - append_next_token_ids = torch.ones( - size=(token_ids.shape[0], self.backend.mtp_step - draft_step), - dtype=token_ids.dtype, - device=token_ids.device, - ) - return torch.cat([token_ids, append_next_token_ids], dim=-1) - - def scatter_next_tokens( - self, - *, - b_req_mtp_start_loc: torch.Tensor, - all_next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - mtp_accept_len: torch.Tensor, - all_next_token_probs: Optional[torch.Tensor] = None, - ) -> None: - self.verifier.scatter_next_tokens( - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, - all_next_token_probs=all_next_token_probs, - ) - return - - def scatter_token_id_steps( - self, - *, - token_id_steps: List[torch.Tensor], - b_req_mtp_start_loc: torch.Tensor, - b_req_idx: torch.Tensor, - mtp_accept_len: torch.Tensor, - row_count: int = None, - ) -> torch.Tensor: - """Stack proposal token columns and scatter them for the next verify.""" - - all_next_token_ids = torch.stack(token_id_steps, dim=1) - if row_count is not None: - all_next_token_ids = all_next_token_ids[: int(row_count), :] - self.scatter_next_tokens( - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, - ) - return all_next_token_ids - - def _get_target_layer_ids(self, model) -> List[int]: - if self._target_layer_ids is not None: - return self._target_layer_ids - - draft_model_dir = self.backend.args.mtp_draft_model_dir - if isinstance(draft_model_dir, list): - draft_model_dir = draft_model_dir[0] - - if draft_model_dir: - with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: - draft_config = json.load(json_file) - normalize_speculative_draft_config(draft_config) - target_layer_ids = draft_config.get("target_layer_ids") - if target_layer_ids is not None: - self._target_layer_ids = [int(layer_id) for layer_id in target_layer_ids] - return self._target_layer_ids - - self._target_layer_ids = [1, model.config["n_layer"] // 2 - 1, model.config["n_layer"] - 4] - return self._target_layer_ids - - def _build_decode_planner(self): - if self.enable_dynamic_mtp: - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - assert g_infer_context.dynamic_mtp_planner is not None - return g_infer_context.dynamic_mtp_planner - return FixedMTPPlanner(self.backend.mtp_step) - - def _clear_stale_dynamic_token_probs(self, *, pre_draft_step: int) -> None: - # Columns after the previous draft length are stale and must not be - # sampled by dynamic row compaction in the current target forward. - self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs[ - :, (pre_draft_step + 1) : - ].fill_(0.0) - return - - -def build_spec_runtime(backend) -> SpecRuntime: - return SpecRuntime(backend) diff --git a/lightllm/server/router/model_infer/speculative/state.py b/lightllm/server/router/model_infer/speculative/state.py deleted file mode 100644 index af16ce8510..0000000000 --- a/lightllm/server/router/model_infer/speculative/state.py +++ /dev/null @@ -1,130 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, List, Optional - -import torch - -from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor - -if TYPE_CHECKING: - from lightllm.server.router.model_infer.speculative.runtime import SpecRuntime - - -class SpecHiddenStore: - """Stores target-model features that are passed into draft-model forwards. - - LightLLM's service path does not materialize HuggingFace-style - `hidden_states`. Instead, BaseModel calls SpecForwardContext while the - target model is running. - - Captured tensors are keyed by microbatch index: - - vanilla MTP consumes the final target hidden state: - [token_num, hidden_size] - - Eagle3 / DSpark-style draft models consume selected target layers - concatenated on the hidden dimension after TP/SP all-gather: - [token_num, hidden_size * len(target_layer_ids)] - - Draft-model forwards are identified by `mtp_draft_input_hiddens is not - None`; their final hidden can also be captured because chained MTP drafts - pass one draft's hidden into the next draft model. - """ - - def __init__(self, runtime: "SpecRuntime") -> None: - self.runtime = runtime - self._captured_hiddens = {} - - def select_hidden(self, *, infer_state, hidden: torch.Tensor, final_hidden: torch.Tensor) -> torch.Tensor: - if self.runtime.is_draft_forward(infer_state): - return final_hidden - if self.runtime.needs_intermediate_target_hidden: - return hidden - return final_hidden - - def capture_hidden(self, *, infer_state, hidden: torch.Tensor, final_hidden: torch.Tensor) -> torch.Tensor: - selected_hidden = self.select_hidden( - infer_state=infer_state, - hidden=hidden, - final_hidden=final_hidden, - ) - self._captured_hiddens[infer_state.microbatch_index] = selected_hidden - return selected_hidden - - def get_hidden(self, microbatch_index: int = 0) -> torch.Tensor: - hidden = self._captured_hiddens.get(microbatch_index) - assert hidden is not None - return hidden - - def unpad_hidden(self, *, token_num: int, microbatch_index: int = 0) -> None: - hidden = self._captured_hiddens.get(microbatch_index) - if hidden is not None and hidden.shape[0] > token_num: - self._captured_hiddens[microbatch_index] = hidden[0:token_num] - return - - def export_graph_capture(self): - captured_hiddens = { - microbatch_index: tensor_to_no_ref_tensor(hidden) - for microbatch_index, hidden in self._captured_hiddens.items() - } - self._captured_hiddens = dict(captured_hiddens) - return captured_hiddens - - def restore_graph_capture(self, captured_hiddens) -> None: - if captured_hiddens is not None: - self._captured_hiddens = dict(captured_hiddens) - return - - -class SpecForwardContext: - """Per-forward hidden capture context used by BaseModel. - - The target model and the draft model exchange only tensors, not model - outputs. During target forward, BaseModel calls `add_hidden` after each - transformer layer. The runtime decides which layer ids matter for the - active draft algorithm. At the end of forward, BaseModel calls `capture` - with: - - `hidden`: selected intermediate target feature, shape - [token_num, hidden_size * selected_layer_num] after all-gather - - `final_hidden`: final target feature, shape [token_num, hidden_size] - - Vanilla MTP uses `final_hidden`; Eagle3/DSpark-style proposers use - `hidden`. - """ - - def __init__(self, *, runtime: "SpecRuntime", model, infer_state) -> None: - self.runtime = runtime - self.model = model - self.infer_state = infer_state - self.layer_ids = runtime.get_capture_layer_ids(model, infer_state) - self.layer_hiddens: List[torch.Tensor] = [] - - def add_hidden(self, *, layer_index: int, layer_num: int, hidden: torch.Tensor) -> None: - if layer_index not in self.layer_ids: - return - - if layer_index == layer_num - 1: - self.layer_hiddens.append(hidden) - else: - self.layer_hiddens.append(hidden.clone()) - return - - def build_layer_hidden(self) -> Optional[torch.Tensor]: - if not self.layer_hiddens: - return None - if len(self.layer_hiddens) == 1: - return self.layer_hiddens[0] - return torch.cat(self.layer_hiddens, dim=-1) - - def capture(self, *, hidden: torch.Tensor, final_hidden: torch.Tensor) -> torch.Tensor: - return self.runtime.capture_hidden( - infer_state=self.infer_state, - hidden=hidden, - final_hidden=final_hidden, - ) - - def capture_final_hidden(self, final_hidden: torch.Tensor) -> None: - assert not self.layer_ids, ( - f"{self.runtime.spec_config.mode} needs intermediate hidden layers and does not support " - "microbatch overlap forward now" - ) - self.capture(hidden=final_hidden, final_hidden=final_hidden) - return diff --git a/lightllm/server/router/model_infer/speculative/verifier.py b/lightllm/server/router/model_infer/speculative/verifier.py index 1a29db0ff5..0272ee346b 100644 --- a/lightllm/server/router/model_infer/speculative/verifier.py +++ b/lightllm/server/router/model_infer/speculative/verifier.py @@ -28,7 +28,7 @@ class SpecVerifyResult: class SpecVerifier: - """Service verifier for LightLLM MTP layout. + """Service verifier for LightLLM's speculative row layout. DeepSpec verifies by running target logits over [current_token + draft_tokens] and applying rejection sampling. LightLLM's @@ -44,7 +44,6 @@ def __init__(self, backend) -> None: def verify_target_tokens( self, - *, new_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, @@ -55,35 +54,34 @@ def verify_target_tokens( - `new_next_token_ids`: target sampled ids, shape [verify_batch] - `b_req_idx`: request ids for each target row, shape [verify_batch] - `b_req_mtp_start_loc`: first row of each logical request in the - MTP-expanded target batch, shape [logical_req_num] + speculative verify batch, shape [logical_req_num] """ - mtp_accept_len, accepted_index = mtp_verify( + spec_accept_len, accepted_index = mtp_verify( req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, new_next_token_ids=new_next_token_ids, b_req_idx=b_req_idx, ) - return SpecVerifyResult(accept_len=mtp_accept_len, accepted_index=accepted_index) + return SpecVerifyResult(accept_len=spec_accept_len, accepted_index=accepted_index) def scatter_next_tokens( self, - *, b_req_mtp_start_loc: torch.Tensor, all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, - mtp_accept_len: torch.Tensor, + spec_accept_len: torch.Tensor, all_next_token_probs: Optional[torch.Tensor] = None, ) -> None: """Scatter target+draft candidates into per-request next-token buffers. Inputs: - - `all_next_token_ids`: [verify_batch, mtp_step + 1]. Column 0 is the + - `all_next_token_ids`: [verify_batch, max_draft_step + 1]. Column 0 is the target sampled token from this iteration; remaining columns are draft - candidates padded to the static MTP width when dynamic MTP produces a - shorter proposal. - - `all_next_token_probs`: optional [verify_batch, mtp_step + 1]. - Current dynamic MTP stores selected-token probabilities, not full + candidates padded to the configured speculative width when dynamic scheduling + produces a shorter proposal. + - `all_next_token_probs`: optional [verify_batch, max_draft_step + 1]. + Dynamic scheduling stores selected-token probabilities, not full vocab distributions. """ @@ -92,11 +90,7 @@ def scatter_next_tokens( b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=all_next_token_ids, b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, - # A profitable full-width Eagle3 plan deliberately uses Static's - # argmax-only proposer and therefore has no confidence matrix to - # scatter. Leave the probability buffer untouched in that mode; - # the next observe/confidence plan refreshes it before compaction. + spec_accept_len=spec_accept_len, req_to_next_token_probs=( self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs if all_next_token_probs is not None @@ -104,4 +98,3 @@ def scatter_next_tokens( ), all_next_token_probs=all_next_token_probs, ) - return diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index b16e1612ac..30f8bea52e 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -5,7 +5,6 @@ from easydict import EasyDict from functools import lru_cache from lightllm.utils.log_utils import init_logger -from lightllm.common.speculative import SpeculativeConfig logger = init_logger(__name__) @@ -223,14 +222,19 @@ def enable_diverse_mode_gqa_decode_fast_kernel() -> bool: @lru_cache(maxsize=None) -def enable_dynamic_mtp_verify() -> bool: - """ - 启用动态 MTP 长度验证功能 - 在 MTP 模式下,根据每步的 prob 分布动态调整验证长度 - 通过启动参数 --mtp_dynamic_verify 控制;DSpark 模式固定使用 - confidence-scheduled dynamic verify。 +def enable_dynamic_spec() -> bool: + """Whether speculative scheduling may vary draft and verify widths. + + ``--mtp_dynamic_verify`` remains the compatible command-line switch; + DSpark enables dynamic speculative scheduling unconditionally. """ - return SpeculativeConfig.from_args(get_env_start_args()).dynamic_verify + + args = get_env_start_args() + if args.mtp_mode == "dspark": + return True + if args.mtp_mode == "dflash": + return False + return bool(args.mtp_dynamic_verify) @lru_cache(maxsize=None) @@ -267,20 +271,20 @@ def enable_cpu_cache_numa_interleave() -> bool: @lru_cache(maxsize=None) def get_added_mtp_kv_layer_num() -> int: + # mtp 模式下需要在mem manger上扩展draft model使用的layer args = get_env_start_args() - spec_config = SpeculativeConfig.from_args(args) - if not spec_config.uses_attention_draft: + spec_mode = args.mtp_mode + if spec_mode not in ("vanilla_with_att", "eagle_with_att", "eagle3", "dspark", "dflash"): return 0 - if spec_config.is_dflash or spec_config.is_dspark or spec_config.is_eagle3: - draft_model_dir = args.mtp_draft_model_dir - if isinstance(draft_model_dir, list): - draft_model_dir = draft_model_dir[0] - if not draft_model_dir: - return spec_config.draft_model_count + draft_model_count = args.mtp_step if spec_mode == "vanilla_with_att" else 1 + if spec_mode in ("dflash", "dspark", "eagle3"): + if not args.mtp_draft_model_dir: + return draft_model_count + draft_model_dir = args.mtp_draft_model_dir[0] with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: draft_config = json.load(json_file) - return int(draft_config.get("num_hidden_layers", draft_config.get("n_layer", spec_config.draft_model_count))) - return spec_config.draft_model_count + return int(draft_config.get("num_hidden_layers", draft_config.get("n_layer", draft_model_count))) + return draft_model_count @lru_cache(maxsize=None) diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index 6856382996..e81caafe7a 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -16,17 +16,17 @@ get_added_mtp_kv_layer_num, ) from lightllm.utils.log_utils import init_logger -from lightllm.common.speculative import SpeculativeConfig from lightllm.utils.config_utils import get_num_key_value_heads, get_head_dim, get_layer_num, is_linear_att_mixed_model from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class from lightllm.common.kv_cache_mem_manager import ( MemoryManager, PPLINT8KVMemoryManager, + PPLINT4KVMemoryManager, Deepseek2MemoryManager, Qwen3NextMemManager, ) -from typing import List, Tuple +from typing import List, Tuple, Optional from tqdm import tqdm from lightllm.utils.auto_shm_cleanup import register_sysv_shm_for_cleanup from lightllm.utils.dist_utils import get_current_device_id @@ -119,10 +119,13 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": logger.error(f"not support mem manager: {mem_manager_class} for cpu kv cache") raise Exception(f"not support mem manager: {mem_manager_class} for cpu kv cache") - spec_config = SpeculativeConfig.from_args(args) - if spec_config.enabled and mem_manager_class is not Qwen3NextMemManager: + if args.mtp_mode is not None: # TODO 可能会存在不同mtp模式的精度问题 - cpu_cache_meta.layer_num += get_added_mtp_kv_layer_num() + if not is_linear_att_mixed_model(args.model_dir): + # 对于非 linear att 混合模型,需要额外增加 mtp 的 kv 层数, + # 对于 linear att 混合模型,如qwen 3.5 mtp,已经将 kv 数据 + # 打包成一个块了,所以不需要额外增加,其 layer_num 一直都保持为 1 + cpu_cache_meta.layer_num += get_added_mtp_kv_layer_num() cpu_cache_page_num = int( (args.cpu_cache_storage_size * 1024 * 1024 * 1024) / (cpu_cache_meta.calcu_one_page_size()) diff --git a/lightllm/utils/sgl_utils.py b/lightllm/utils/sgl_utils.py index 9bacece6a8..c39e8bc142 100644 --- a/lightllm/utils/sgl_utils.py +++ b/lightllm/utils/sgl_utils.py @@ -122,7 +122,7 @@ def flash_attn_with_kvcache_autotune( ) -def fa3_decode_autotune(model, cuda_graph_batch_sizes): +def fa3_decode_autotune(model, cuda_graph_batch_sizes, batch_multiplier: int): # 是否开启自动调优 if get_triton_autotune_level() not in [ AutotuneLevel.ADAPTIVE_AUTOTUNE, @@ -141,12 +141,9 @@ def fa3_decode_autotune(model, cuda_graph_batch_sizes): v_cache = v.view(v.shape[0], 1, v.shape[1], v.shape[2]) q_head_num = int(model.config["num_attention_heads"]) // model.tp_world_size_ head_dim = int(k.shape[-1]) - mtp_size = model.args.mtp_step + 1 - for batch_size in cuda_graph_batch_sizes[::-1]: - att_batch_size = batch_size // mtp_size - if att_batch_size <= 0: - continue + assert batch_size % batch_multiplier == 0 + att_batch_size = batch_size // batch_multiplier # 因为完整的kv空间可能无法装下所有token,所以在tuning的时候,所有token都使用相同的kv空间。 # 保证tuning的时候不会出现大的问题。 kv_range = torch.arange(att_batch_size * max_kv_len, dtype=torch.int32, device=k.device) % max_kv_len @@ -154,13 +151,13 @@ def fa3_decode_autotune(model, cuda_graph_batch_sizes): v[kv_range].zero_() q = torch.zeros( - (att_batch_size * mtp_size, q_head_num, head_dim), + (batch_size, q_head_num, head_dim), dtype=model.data_type, device=k.device, ) page_table = kv_range.view(att_batch_size, max_kv_len) cache_seqlens = torch.full((att_batch_size,), max_kv_len, dtype=torch.int32, device=k.device) - cu_seqlens_q = torch.arange(att_batch_size + 1, dtype=torch.int32, device=k.device) * mtp_size + cu_seqlens_q = torch.arange(att_batch_size + 1, dtype=torch.int32, device=k.device) * batch_multiplier cu_seqlens_k = torch.arange(att_batch_size + 1, dtype=torch.int32, device=k.device) * max_kv_len softmax_scale = 1.0 / (head_dim ** 0.5) @@ -172,7 +169,7 @@ def fa3_decode_autotune(model, cuda_graph_batch_sizes): cache_seqlens=cache_seqlens, cu_seqlens_q=cu_seqlens_q, cu_seqlens_k_new=cu_seqlens_k, - max_seqlen_q=mtp_size, + max_seqlen_q=batch_multiplier, softmax_scale=softmax_scale, causal=True, window_size=(-1, -1), diff --git a/test/speculative/test_qwen35_dflash_state.py b/test/speculative/test_qwen35_dflash_state.py deleted file mode 100644 index eed481e683..0000000000 --- a/test/speculative/test_qwen35_dflash_state.py +++ /dev/null @@ -1,378 +0,0 @@ -from types import MethodType, SimpleNamespace - -import pytest -import torch - -from lightllm.common.speculative import BlockDraftLayout, SpeculativeConfig, get_block_draft_layout -from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel -from lightllm.models.qwen3_dflash.layer_infer.transformer_layer_infer import ( - get_draft_layer_type, -) -from lightllm.common.kv_cache_mem_manager.qwen3next_mem_manager import Qwen3NextMemManager -from lightllm.server import api_start -from lightllm.utils import envs_utils, kv_cache_utils -from lightllm.server.router.model_infer.mode_backend.pd.prefill_node_impl.prefill_impl import ( - PDChunkedPrefillForPrefillNode, -) -from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend -from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer - - -def test_block_layout_skips_leading_bonus_query_prediction(): - block_token_ids = torch.arange(32).reshape(2, 16) - - selected = DFlashProposer.select_draft_token_ids( - block_token_ids=block_token_ids, - draft_step=15, - layout=BlockDraftLayout(query_block_size=16, proposal_output_start=1), - ) - - torch.testing.assert_close(selected, block_token_ids[:, 1:]) - - -def test_block_layout_keeps_first_query_prediction(): - block_token_ids = torch.arange(32).reshape(2, 16) - - selected = DFlashProposer.select_draft_token_ids( - block_token_ids=block_token_ids, - draft_step=16, - layout=BlockDraftLayout(query_block_size=16, proposal_output_start=0), - ) - - torch.testing.assert_close(selected, block_token_ids) - - -def test_block_draft_layout_defaults_to_zero_and_uses_architecture_override(): - draft_fields = { - "block_size": 16, - "target_layer_ids": [1, 10], - "mask_token_id": 248077, - } - nested = { - "architectures": ["Qwen3DFlashModel"], - "dflash_config": draft_fields, - } - flat = {"architectures": ["Qwen3DFlashModel"], **draft_fields, "block_size": 7} - qwen35 = { - "architectures": ["Qwen3_5DFlashModel"], - "dflash_config": draft_fields, - } - - assert get_block_draft_layout( - nested, - mode="dflash", - ) == BlockDraftLayout(query_block_size=16, proposal_output_start=0) - assert get_block_draft_layout( - flat, - mode="dflash", - ) == BlockDraftLayout(query_block_size=7, proposal_output_start=0) - assert get_block_draft_layout(qwen35, mode="dflash") == BlockDraftLayout( - query_block_size=16, proposal_output_start=1 - ) - - -def test_qwen35_dflash_normalizes_block_size_to_fifteen_draft_tokens(): - backend = ModeBackend.__new__(ModeBackend) - backend.args = SimpleNamespace(mtp_step=16) - backend.mtp_step = 16 - backend.spec_config = SpeculativeConfig(mode="dflash", step=16) - backend.logger = SimpleNamespace(warning=lambda *args: None) - config = { - "architectures": ["Qwen3_5DFlashModel"], - "dflash_config": { - "block_size": 16, - "target_layer_ids": [1, 10], - "mask_token_id": 248077, - }, - } - - backend._normalize_block_mtp_step_from_config(config) - - assert backend.args.mtp_step == 15 - assert backend.mtp_step == 15 - assert backend.spec_config.step == 15 - assert backend.block_draft_layout == BlockDraftLayout(query_block_size=16, proposal_output_start=1) - - -def test_qwen35_dflash_respects_shorter_configured_draft_step(): - backend = ModeBackend.__new__(ModeBackend) - backend.args = SimpleNamespace(mtp_step=9) - backend.mtp_step = 9 - backend.spec_config = SpeculativeConfig(mode="dflash", step=9) - backend.logger = SimpleNamespace(warning=lambda *args: None) - config = { - "architectures": ["Qwen3_5DFlashModel"], - "dflash_config": { - "block_size": 16, - "target_layer_ids": [1, 10], - "mask_token_id": 248077, - }, - } - - backend._normalize_block_mtp_step_from_config(config) - - assert backend.args.mtp_step == 9 - assert backend.mtp_step == 9 - assert backend.spec_config.step == 9 - - -def test_qwen35_dflash_normalizes_global_start_args_to_fifteen(monkeypatch): - config = { - "architectures": ["Qwen3_5DFlashModel"], - "dflash_config": { - "block_size": 16, - "target_layer_ids": [1, 10], - "mask_token_id": 248077, - }, - } - monkeypatch.setattr( - api_start.PretrainedConfig, - "get_config_dict", - lambda *_args, **_kwargs: (config, {}), - ) - args = SimpleNamespace(mtp_step=16, mtp_draft_model_dir=["unused"]) - - spec_config = api_start.normalize_block_mtp_step_from_first_draft_config( - args, - SpeculativeConfig(mode="dflash", step=16), - ) - - assert args.mtp_step == 15 - assert spec_config.step == 15 - - -def test_qwen35_dflash_global_normalization_respects_shorter_step(monkeypatch): - config = { - "architectures": ["Qwen3_5DFlashModel"], - "dflash_config": { - "block_size": 16, - "target_layer_ids": [1, 10], - "mask_token_id": 248077, - }, - } - monkeypatch.setattr( - api_start.PretrainedConfig, - "get_config_dict", - lambda *_args, **_kwargs: (config, {}), - ) - args = SimpleNamespace(mtp_step=9, mtp_draft_model_dir=["unused"]) - - spec_config = api_start.normalize_block_mtp_step_from_first_draft_config( - args, - SpeculativeConfig(mode="dflash", step=9), - ) - - assert args.mtp_step == 9 - assert spec_config.step == 9 - - -def test_qwen35_dflash_support_scope_preserves_existing_lightspec_modes(monkeypatch): - backend = ModeBackend.__new__(ModeBackend) - backend.is_linear_att_mixed_model = True - backend.args = SimpleNamespace(mtp_draft_model_dir=["unused"]) - - backend.spec_config = SpeculativeConfig(mode="qwen3next_eagle", step=2, dynamic_verify=True) - backend._validate_linear_att_spec_support() - - backend.spec_config = SpeculativeConfig(mode="dflash", step=6) - monkeypatch.setattr( - "lightllm.server.router.model_infer.mode_backend.base_backend.PretrainedConfig.get_config_dict", - lambda *_args, **_kwargs: ({"architectures": ["Qwen3DFlashModel"]}, {}), - ) - with pytest.raises(AssertionError, match="Qwen3_5DFlashModel"): - backend._validate_linear_att_spec_support() - - monkeypatch.setattr( - "lightllm.server.router.model_infer.mode_backend.base_backend.PretrainedConfig.get_config_dict", - lambda *_args, **_kwargs: ({"architectures": ["Qwen3_5DFlashModel"]}, {}), - ) - backend.is_linear_att_mixed_model = False - with pytest.raises(AssertionError, match="requires a Qwen3Next target"): - backend._validate_linear_att_spec_support() - - backend.is_linear_att_mixed_model = True - backend._validate_linear_att_spec_support() - - - -def test_qwen35_dflash_restores_target_allocator_on_shared_req_manager(): - target_mem_manager = object() - draft_mem_manager = object() - req_manager = SimpleNamespace(mem_manager=draft_mem_manager) - model = Qwen3_5DFlashModel.__new__(Qwen3_5DFlashModel) - model.main_model = SimpleNamespace(req_manager=req_manager, mem_manager=target_mem_manager) - model.req_manager = req_manager - model.mem_manager = draft_mem_manager - - model._restore_main_mem_manager() - - assert model.req_manager.mem_manager is target_mem_manager - assert model.mem_manager is draft_mem_manager - - -def test_qwen35_pd_move_page_supports_distinct_target_and_draft_kv_shapes(): - manager = Qwen3NextMemManager.__new__(Qwen3NextMemManager) - manager.size = 16 - manager.dtype = torch.float32 - manager.layer_num = 2 - manager.head_dim = 4 - manager.linear_config = SimpleNamespace(full_att_all_num_kv_heads=3) - manager._pd_dflash_draft_mem_manager = None - manager._pd_dflash_global_kv_heads = None - draft_manager = SimpleNamespace( - size=16, - dtype=torch.float32, - layer_num=5, - head_dim=2, - ) - - manager.register_dflash_draft_mem_manager(draft_manager, global_kv_heads=2) - elements_per_token = max(2 * 2 * 3 * 4, 5 * 2 * 2 * 2) - manager.kv_move_buffer = torch.empty((1, 7, 1, 1, elements_per_token)) - - target_page = manager._get_pd_kv_page(0, "kv") - draft_page = manager._get_pd_kv_page(0, "draft_kv") - assert target_page.shape == (7, 2, 6, 4) - assert draft_page.shape == (7, 5, 4, 2) - assert target_page.is_contiguous() - assert draft_page.is_contiguous() - - -def test_qwen35_pd_prefill_emits_target_and_draft_kv_tasks(): - emitted_page_kinds = [] - - backend = PDChunkedPrefillForPrefillNode.__new__(PDChunkedPrefillForPrefillNode) - backend.args = SimpleNamespace(pd_kv_page_size=4) - backend.model = SimpleNamespace(mem_manager=SimpleNamespace(has_separate_dflash_draft_kv=True)) - backend.is_master_in_dp = False - - def fake_create_task(self, req_obj, kv_start_index, kv_end_index, page_kind="kv"): - del self, req_obj, kv_start_index, kv_end_index - emitted_page_kinds.append(page_kind) - return SimpleNamespace(first_gen_token_id=None, first_gen_token_logprob=None) - - backend._create_pd_trans_task = MethodType(fake_create_task, backend) - req = SimpleNamespace( - cur_kv_len=4, - pd_trans_kv_start_index=0, - shm_req=SimpleNamespace(input_len=4), - ) - - backend._prefill_chuncked_handle_func(req, next_token_id=1, next_token_prob=0.0, output_len=0) - - assert emitted_page_kinds == ["kv", "draft_kv"] - assert req.pd_trans_kv_start_index == 4 - - -def test_dflash_layer_type_uses_draft_local_index_for_global_layer_numbers(): - config = { - "_draft_layer_start": 64, - "layer_types": ["sliding_attention"] * 5 + ["full_attention"], - } - - assert [get_draft_layer_type(i, config) for i in range(64, 70)] == [ - "sliding_attention", - "sliding_attention", - "sliding_attention", - "sliding_attention", - "sliding_attention", - "full_attention", - ] - - - -def test_qwen35_cpu_cache_appends_distinct_draft_kv_region(monkeypatch): - page_num = 2 - page_size = 4 - target_bytes = 32 - manager = Qwen3NextMemManager.__new__(Qwen3NextMemManager) - manager.size = 16 - manager.dtype = torch.float32 - manager.linear_config = SimpleNamespace(get_cpu_cache_big_page_bytes=lambda: target_bytes) - manager._pd_dflash_draft_mem_manager = None - manager._pd_dflash_global_kv_heads = None - draft_manager = SimpleNamespace( - size=16, - dtype=torch.float32, - layer_num=3, - head_dim=2, - ) - manager.register_dflash_draft_mem_manager(draft_manager, global_kv_heads=2) - monkeypatch.setattr( - "lightllm.common.kv_cache_mem_manager.qwen3next_mem_manager.get_env_start_args", - lambda: SimpleNamespace(cpu_cache_token_page_size=page_size), - ) - - draft_bytes = 3 * page_size * 4 * 2 * torch.float32.itemsize - cpu_cache = torch.full((page_num, 1, 1, 1, target_bytes + draft_bytes), 0xA5, dtype=torch.uint8) - draft_cache = manager.get_dflash_draft_cpu_cache(cpu_cache) - - assert draft_cache.shape == (page_num, 3, page_size, 4, 2) - assert draft_cache.dtype == torch.float32 - draft_cache.zero_() - assert torch.all(cpu_cache.reshape(page_num, -1)[:, :target_bytes] == 0xA5) - assert torch.all(cpu_cache.reshape(page_num, -1)[:, target_bytes:] == 0) - - -def test_qwen35_dflash_does_not_reserve_duplicate_target_kv_layers(tmp_path, monkeypatch): - (tmp_path / "config.json").write_text( - '{"architectures": ["Qwen3_5DFlashModel"], "n_layer": 5}', - encoding="utf-8", - ) - args = SimpleNamespace( - mtp_mode="dflash", - mtp_step=15, - mtp_dynamic_verify=False, - mtp_draft_model_dir=[str(tmp_path)], - ) - monkeypatch.setattr(envs_utils, "get_env_start_args", lambda: args) - envs_utils.get_added_mtp_kv_layer_num.cache_clear() - - assert envs_utils.get_added_mtp_kv_layer_num() == 0 - envs_utils.get_added_mtp_kv_layer_num.cache_clear() - - -def test_qwen35_dflash_cpu_cache_meta_includes_draft_kv(monkeypatch): - args = SimpleNamespace( - enable_cpu_cache=True, - model_dir="unused-target", - mtp_mode="dflash", - mtp_step=15, - mtp_dynamic_verify=False, - cpu_cache_storage_size=1, - ) - monkeypatch.setattr(kv_cache_utils, "get_env_start_args", lambda: args) - monkeypatch.setattr(kv_cache_utils, "is_linear_att_mixed_model", lambda _path: True) - monkeypatch.setattr(kv_cache_utils, "get_llm_data_type", lambda: torch.bfloat16) - monkeypatch.setattr( - kv_cache_utils.LinearAttCacheConfig, - "load_from_args", - lambda: SimpleNamespace(get_cpu_cache_big_page_bytes=lambda: 1024), - ) - monkeypatch.setattr(kv_cache_utils, "_get_qwen35_dflash_cpu_cache_bytes", lambda _args: 384) - kv_cache_utils.calcu_cpu_cache_meta.cache_clear() - - meta = kv_cache_utils.calcu_cpu_cache_meta() - - assert meta.data_type == torch.uint8 - assert meta.calcu_one_page_size() == 1024 + 384 - kv_cache_utils.calcu_cpu_cache_meta.cache_clear() - - -def test_qwen35_dflash_cpu_cache_draft_size_uses_checkpoint_shape(monkeypatch): - from transformers.configuration_utils import PretrainedConfig - - monkeypatch.setattr( - PretrainedConfig, - "get_config_dict", - lambda *_args, **_kwargs: ({"architectures": ["Qwen3_5DFlashModel"]}, {}), - ) - monkeypatch.setattr(kv_cache_utils, "get_layer_num", lambda _path: 3) - monkeypatch.setattr(kv_cache_utils, "get_num_key_value_heads", lambda _path: 2) - monkeypatch.setattr(kv_cache_utils, "get_head_dim", lambda _path: 5) - monkeypatch.setattr(kv_cache_utils, "get_llm_data_type", lambda: torch.bfloat16) - args = SimpleNamespace(mtp_draft_model_dir=["unused-draft"], cpu_cache_token_page_size=4) - - draft_bytes = kv_cache_utils._get_qwen35_dflash_cpu_cache_bytes(args) - - assert draft_bytes == 4 * 3 * 2 * 2 * 5 * torch.bfloat16.itemsize diff --git a/unit_tests/common/basemodel/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py new file mode 100644 index 0000000000..1ba29894ef --- /dev/null +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -0,0 +1,71 @@ +from types import SimpleNamespace + +import pytest + +import lightllm.common.basemodel.cuda_graph as cuda_graph_module +from lightllm.common.basemodel.cuda_graph import CudaGraph + + +@pytest.fixture(autouse=True) +def _graph_args(monkeypatch): + args = SimpleNamespace( + graph_split_batch_size=4, + graph_grow_step_size=2, + enable_decode_microbatch_overlap=False, + enable_tpsp_mix_mode=False, + enable_torch_memory_saver=False, + ) + monkeypatch.setattr(cuda_graph_module, "get_env_start_args", lambda: args) + return args + + +def _batch_sizes(max_batch_size, batch_multiplier=1): + physical_max_batch_size = max_batch_size * batch_multiplier + graph = CudaGraph( + max_batch_size=physical_max_batch_size, + batch_multiplier=batch_multiplier, + ) + return graph.cuda_graph_batch_sizes + + +def test_dynamic_schedule_uses_compacted_physical_rows(_graph_args): + assert _batch_sizes(max_batch_size=128) == [1, 2, 3, 4, *range(6, 129, 2)] + + +def test_public_static_schedule_preserves_original_static_mtp_default(_graph_args): + assert CudaGraph.gen_cuda_graph_batch_sizes(max_batch_size=32, batch_multiplier=8) == [ + 8, + 16, + 24, + 32, + ] + + +def test_instance_and_public_static_schedule_match(_graph_args): + graph = CudaGraph(max_batch_size=128, batch_multiplier=8) + + assert graph.cuda_graph_batch_sizes == CudaGraph.gen_cuda_graph_batch_sizes( + max_batch_size=graph.max_batch_size, + tp_world_size=graph.tp_world_size, + batch_multiplier=8, + ) + + +def test_legacy_vanilla_layout_can_keep_k_plus_one_stride(_graph_args): + assert _batch_sizes(max_batch_size=4, batch_multiplier=8) == [ + 8, + 16, + 24, + 32, + ] + + +def test_block_draft_layout_only_pads_to_complete_physical_blocks(_graph_args): + assert _batch_sizes(max_batch_size=8, batch_multiplier=7) == [ + 7, + 14, + 21, + 28, + 42, + 56, + ] diff --git a/unit_tests/common/basemodel/test_hidden_collector.py b/unit_tests/common/basemodel/test_hidden_collector.py new file mode 100644 index 0000000000..dbd311ee60 --- /dev/null +++ b/unit_tests/common/basemodel/test_hidden_collector.py @@ -0,0 +1,149 @@ +from types import SimpleNamespace + +import torch + +from lightllm.common.basemodel.hidden_collector import ( + FinalHiddenCollector, + HiddenCollector, + LayerHiddenCollector, + NoopHiddenCollector, +) + + +class _IdentityPreInfer: + @staticmethod + def _tpsp_allgather(input, infer_state): + del infer_state + return input + + +def test_hidden_collector_selects_implementation(): + model = SimpleNamespace(is_mtp_draft_model=False, layers_num=3, pre_infer=_IdentityPreInfer()) + noop_collector = HiddenCollector() + final_hidden_collector = HiddenCollector(model=model, spec_mode="vanilla_with_att") + layer_hidden_collector = HiddenCollector( + model=model, + spec_mode="eagle3", + layer_ids=[0, 2], + microbatch_count=2, + ) + + assert isinstance(noop_collector, HiddenCollector) + assert isinstance(noop_collector.collectors[0], NoopHiddenCollector) + assert isinstance(final_hidden_collector.collectors[0], FinalHiddenCollector) + assert all(isinstance(collector, LayerHiddenCollector) for collector in layer_hidden_collector.collectors) + + +def test_draft_hidden_collector_follows_spec_mode(): + model = SimpleNamespace(is_mtp_draft_model=True) + + recurrent_collector = HiddenCollector(model=model, spec_mode="eagle3") + block_collector = HiddenCollector(model=model, spec_mode="dspark") + + assert isinstance(recurrent_collector.collectors[0], FinalHiddenCollector) + assert isinstance(block_collector.collectors[0], NoopHiddenCollector) + + +def test_hidden_collector_supports_single_and_overlap_forward(): + model = SimpleNamespace(is_mtp_draft_model=False) + collector = HiddenCollector(model=model, spec_mode="vanilla_with_att", microbatch_count=2) + hidden0 = torch.randn(2, 3) + hidden1 = torch.randn(2, 3) + infer_state = SimpleNamespace(need_dp_prefill_balance=False) + + collector.add(layer_index=0, hidden=hidden0) + collected = collector.finish( + infer_state=infer_state, + final_hidden=hidden0, + ) + assert collected.data_ptr() == hidden0.data_ptr() + + collector.add(layer_index=0, hidden=hidden0) + collector.add(layer_index=0, hidden=hidden1, microbatch_index=1) + collected0 = collector.finish( + infer_state=infer_state, + final_hidden=hidden0, + ) + collected1 = collector.finish( + infer_state=infer_state, + final_hidden=hidden1, + microbatch_index=1, + ) + assert collected0.data_ptr() == hidden0.data_ptr() + assert collected1.data_ptr() == hidden1.data_ptr() + + +def test_layer_hidden_collector_keeps_microbatch_state_separate(): + model = SimpleNamespace(is_mtp_draft_model=False, layers_num=2, pre_infer=_IdentityPreInfer()) + collector = HiddenCollector(model=model, spec_mode="eagle3", layer_ids=[0], microbatch_count=2) + hidden0 = torch.full((2, 3), 1.0) + hidden1 = torch.full((2, 3), 2.0) + infer_state = SimpleNamespace(need_dp_prefill_balance=False) + + collector.add(layer_index=0, hidden=hidden0) + collector.add(layer_index=0, hidden=hidden1, microbatch_index=1) + + collected0 = collector.finish(infer_state=infer_state, final_hidden=hidden0) + collected1 = collector.finish(infer_state=infer_state, final_hidden=hidden1, microbatch_index=1) + + assert torch.equal(collected0, hidden0) + assert torch.equal(collected1, hidden1) + + +def test_noop_collector_keeps_normal_forward_output_minimal(): + final_hidden = torch.randn(2, 3) + collector = NoopHiddenCollector() + + collector.add(layer_index=0, hidden=final_hidden) + + assert collector.prefill_outputs(final_hidden) == [final_hidden] + assert ( + collector.finish( + infer_state=None, + final_hidden=final_hidden, + ) + is None + ) + + +def test_final_collector_returns_final_hidden_without_layer_bookkeeping(): + final_hidden = torch.randn(2, 3) + collected = FinalHiddenCollector().finish( + infer_state=None, + final_hidden=final_hidden, + ) + + assert collected.data_ptr() == final_hidden.data_ptr() + + +def test_layer_collector_preserves_selected_layers_in_model_order(): + layer0 = torch.full((2, 2), 1.0) + layer1 = torch.full((2, 2), 2.0) + layer2 = torch.full((2, 2), 3.0) + model = SimpleNamespace(layers_num=3, pre_infer=_IdentityPreInfer()) + collector = LayerHiddenCollector(model=model, layer_ids=[0, 2]) + + collector.add(layer_index=0, hidden=layer0) + collector.add(layer_index=1, hidden=layer1) + collector.add(layer_index=2, hidden=layer2) + layer0.fill_(9.0) + + forward_outputs = collector.prefill_outputs(layer2) + collected = collector.finish( + infer_state=SimpleNamespace(need_dp_prefill_balance=False), + final_hidden=layer2, + forward_outputs=forward_outputs, + ) + + assert torch.equal(collected, torch.cat([torch.full((2, 2), 1.0), layer2], dim=-1)) + assert not collector.layer_hiddens + + collector.add(layer_index=0, hidden=layer0) + collector.add(layer_index=2, hidden=layer2) + collected = collector.finish( + infer_state=SimpleNamespace(need_dp_prefill_balance=False), + final_hidden=layer2, + ) + + assert torch.equal(collected, torch.cat([layer0, layer2], dim=-1)) + assert not collector.layer_hiddens diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py new file mode 100644 index 0000000000..e66c120123 --- /dev/null +++ b/unit_tests/common/basemodel/test_model_output.py @@ -0,0 +1,40 @@ +import torch + +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.batch_objs import ModelOutput + + +def test_decode_unpad_slices_spec_output_with_logits(): + model = TpPartBaseModel.__new__(TpPartBaseModel) + output = ModelOutput( + logits=torch.arange(24).view(6, 4), + spec_hidden=torch.arange(18).view(6, 3), + ) + + unpadded = model._create_unpad_decode_model_output(output, origin_batch_size=4) + + assert unpadded.logits.shape == (4, 4) + assert unpadded.spec_hidden.shape == (4, 3) + # Unpadding returns a shallow output copy and leaves the graph-owned + # tensors on the original ModelOutput intact. + assert output.logits.shape == (6, 4) + assert output.spec_hidden.shape == (6, 3) + + +def test_prefill_unpad_uses_token_rows_for_spec_hidden(): + model = TpPartBaseModel.__new__(TpPartBaseModel) + output = ModelOutput( + logits=torch.arange(20).view(5, 4), + spec_hidden=torch.arange(24).view(8, 3), + prompt_logics=torch.arange(32).view(8, 4), + ) + + unpadded = model._create_unpad_prefill_model_output( + output, + origin_handle_token_num=6, + origin_batch_size=3, + ) + + assert unpadded.logits.shape == (3, 4) + assert unpadded.spec_hidden.shape == (6, 3) + assert unpadded.prompt_logics.shape == (6, 4) diff --git a/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py b/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py index 816c807c24..3e2555b339 100644 --- a/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py +++ b/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py @@ -54,7 +54,7 @@ def test_token_decode_attention_flash_decoding_diverse_matches_normal_decode(sha ) num_heads = 32 - kv_head_num = 2 # gqa_group_size = 16,满足 Triton tl.dot 的 M >= 16 要求 + kv_head_num = 8 mark_shared_group_size = 3 seq_len = 3547 head_dim = 128 @@ -118,7 +118,6 @@ def test_token_decode_attention_flash_decoding_diverse_matches_normal_decode(sha cache_v_scale=cache_v_scale, alloc_tensor_func=alloc_tensor_func, ) - # 运行 diverse 版本 diverse_out = diverse_attention( q=q.clone(), @@ -130,4 +129,11 @@ def test_token_decode_attention_flash_decoding_diverse_matches_normal_decode(sha alloc_tensor_func=alloc_tensor_func, ) - torch.testing.assert_close(normal_out, diverse_out, atol=1e-2, rtol=1e-2) + print(f"\nshared_seq_len={shared_seq_len}\nbatch_size={batch_size}") + print(f"normal_out: {normal_out[0, 0, :4]}") + print(f"diverse_out: {diverse_out[0, 0, :4]}") + print(f"max diff: {(normal_out - diverse_out).abs().max()}") + + assert torch.allclose( + normal_out, diverse_out, atol=1e-2, rtol=1e-2 + ), f"diverse vs normal decode mismatch for shared_seq_len={shared_seq_len}" diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py similarity index 63% rename from unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py rename to unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py index e99c0ee357..2fd378b7c5 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py @@ -3,36 +3,36 @@ import triton import numpy as np -from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import ( +from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import ( _fwd_kernel_cumprod_probs, - sample_dynamic_mtp_req_mask, + sample_dynamic_spec_row_mask, ) -def _reference_cumprod_probs(req_to_next_token_probs, b_req_idx, mtp_step: int) -> torch.Tensor: +def _reference_cumprod_probs(req_to_next_token_probs, b_req_idx, max_draft_step: int) -> torch.Tensor: probs = req_to_next_token_probs.clone() - req_num = b_req_idx.shape[0] // (mtp_step + 1) + req_num = b_req_idx.shape[0] // (max_draft_step + 1) for req_i in range(req_num): - req_idx = int(b_req_idx[req_i * (mtp_step + 1)].item()) - row = probs[req_idx, : mtp_step + 1].clone() + req_idx = int(b_req_idx[req_i * (max_draft_step + 1)].item()) + row = probs[req_idx, : max_draft_step + 1].clone() row[0] = 1.0 row[1:] = torch.clamp(row[1:], min=0.01, max=0.99) - probs[req_idx, : mtp_step + 1] = torch.cumprod(row, dim=0) + probs[req_idx, : max_draft_step + 1] = torch.cumprod(row, dim=0) return probs def _flat_cumprod_probs( b_req_idx: torch.Tensor, req_to_next_token_probs: torch.Tensor, - mtp_step: int, + max_draft_step: int, ) -> torch.Tensor: - probs = _reference_cumprod_probs(req_to_next_token_probs, b_req_idx, mtp_step) - req_num = b_req_idx.shape[0] // (mtp_step + 1) - all_num = req_num * (mtp_step + 1) + probs = _reference_cumprod_probs(req_to_next_token_probs, b_req_idx, max_draft_step) + req_num = b_req_idx.shape[0] // (max_draft_step + 1) + all_num = req_num * (max_draft_step + 1) flat_probs = [] for offset in range(all_num): req_idx = int(b_req_idx[offset].item()) - mtp_index = offset % (mtp_step + 1) + mtp_index = offset % (max_draft_step + 1) flat_probs.append(probs[req_idx, mtp_index]) return torch.stack(flat_probs) @@ -46,24 +46,24 @@ def _assert_topk_mask(select: torch.Tensor, flat_probs: torch.Tensor, dynamic_ba assert selected_scores.min() >= unselected_scores.max() - 1e-5 -def _make_batch_probs(req_num: int, mtp_step: int, rows): +def _make_batch_probs(req_num: int, max_draft_step: int, rows): max_req = req_num probs = torch.zeros((max_req + 1, 16), dtype=torch.float32, device="cuda") for req_idx, row in enumerate(rows): - probs[req_idx, : mtp_step + 1] = torch.tensor(row, dtype=torch.float32, device="cuda") - b_req_idx = torch.arange(req_num, dtype=torch.int32, device="cuda").repeat_interleave(mtp_step + 1) + probs[req_idx, : max_draft_step + 1] = torch.tensor(row, dtype=torch.float32, device="cuda") + b_req_idx = torch.arange(req_num, dtype=torch.int32, device="cuda").repeat_interleave(max_draft_step + 1) return probs, b_req_idx -@pytest.mark.parametrize("mtp_step", [1, 3]) -def test_cumprod_probs_kernel(mtp_step: int): +@pytest.mark.parametrize("max_draft_step", [1, 3]) +def test_cumprod_probs_kernel(max_draft_step: int): req_num = 2 probs, b_req_idx = _make_batch_probs( req_num, - mtp_step, + max_draft_step, rows=[ - [1.0] + [0.5] * mtp_step, - [1.0] + [0.2] * mtp_step, + [1.0] + [0.5] * max_draft_step, + [1.0] + [0.2] * max_draft_step, ], ) probs_clone = probs.clone() @@ -71,29 +71,29 @@ def test_cumprod_probs_kernel(mtp_step: int): req_to_next_token_probs=probs_clone, req_to_next_token_probs_stride=probs_clone.stride(0), b_req_idx=b_req_idx, - mtp_step=mtp_step, - BLOCK_SIZE=triton.next_power_of_2(mtp_step + 1), + max_draft_step=max_draft_step, + BLOCK_SIZE=triton.next_power_of_2(max_draft_step + 1), num_warps=1, num_stages=1, ) - expected = _reference_cumprod_probs(probs, b_req_idx, mtp_step) - assert torch.allclose(probs_clone[:, : mtp_step + 1], expected[:, : mtp_step + 1], rtol=1e-5, atol=1e-5) + expected = _reference_cumprod_probs(probs, b_req_idx, max_draft_step) + assert torch.allclose(probs_clone[:, : max_draft_step + 1], expected[:, : max_draft_step + 1], rtol=1e-5, atol=1e-5) def test_cumprod_probs_clamps_invalid_values(): - mtp_step = 2 + max_draft_step = 2 req_num = 1 - probs, b_req_idx = _make_batch_probs(req_num, mtp_step, rows=[[1.0, 0.0, 1.5]]) + probs, b_req_idx = _make_batch_probs(req_num, max_draft_step, rows=[[1.0, 0.0, 1.5]]) _fwd_kernel_cumprod_probs[(req_num,)]( req_to_next_token_probs=probs, req_to_next_token_probs_stride=probs.stride(0), b_req_idx=b_req_idx, - mtp_step=mtp_step, - BLOCK_SIZE=triton.next_power_of_2(mtp_step + 1), + max_draft_step=max_draft_step, + BLOCK_SIZE=triton.next_power_of_2(max_draft_step + 1), num_warps=1, num_stages=1, ) - row = probs[0, : mtp_step + 1] + row = probs[0, : max_draft_step + 1] # Index 0 is the always-accepted target sample; clamping applies only to drafts. assert row[0].item() == pytest.approx(1.0) assert row[1].item() == pytest.approx(0.01, rel=1e-4) @@ -101,22 +101,22 @@ def test_cumprod_probs_clamps_invalid_values(): def test_cumprod_probs_clamps_boundary_values(): - mtp_step = 3 + max_draft_step = 3 req_num = 1 - probs, b_req_idx = _make_batch_probs(req_num, mtp_step, rows=[[1.0, 0.995, 0.005, 0.5]]) + probs, b_req_idx = _make_batch_probs(req_num, max_draft_step, rows=[[1.0, 0.995, 0.005, 0.5]]) raw_probs = probs.clone() _fwd_kernel_cumprod_probs[(req_num,)]( req_to_next_token_probs=probs, req_to_next_token_probs_stride=probs.stride(0), b_req_idx=b_req_idx, - mtp_step=mtp_step, - BLOCK_SIZE=triton.next_power_of_2(mtp_step + 1), + max_draft_step=max_draft_step, + BLOCK_SIZE=triton.next_power_of_2(max_draft_step + 1), num_warps=1, num_stages=1, ) - expected = _reference_cumprod_probs(raw_probs, b_req_idx, mtp_step) - row = probs[0, : mtp_step + 1] - assert torch.allclose(row, expected[0, : mtp_step + 1], rtol=1e-5, atol=1e-5) + expected = _reference_cumprod_probs(raw_probs, b_req_idx, max_draft_step) + row = probs[0, : max_draft_step + 1] + assert torch.allclose(row, expected[0, : max_draft_step + 1], rtol=1e-5, atol=1e-5) # Draft probabilities 0.995 and 0.005 clamp to 0.99 and 0.01. assert row[0].item() == pytest.approx(1.0) assert row[1].item() == pytest.approx(0.99, rel=1e-4) @@ -125,24 +125,24 @@ def test_cumprod_probs_clamps_boundary_values(): def test_sample_select_count(): - mtp_step = 3 + max_draft_step = 3 req_num = 3 probs, b_req_idx = _make_batch_probs( req_num, - mtp_step, + max_draft_step, rows=[ [1.0, 0.95, 0.90, 0.10], [1.0, 0.20, 0.80, 0.80], [1.0, 0.99, 0.99, 0.99], ], ) - all_num = req_num * (mtp_step + 1) + all_num = req_num * (max_draft_step + 1) for dynamic_batch_size in [3, 8, all_num]: - select = sample_dynamic_mtp_req_mask( + select = sample_dynamic_spec_row_mask( dynamic_batch_size=dynamic_batch_size, b_req_idx=b_req_idx, req_to_next_token_probs=probs.clone(), - mtp_step=mtp_step, + max_draft_step=max_draft_step, ) assert select.dtype == torch.int32 assert select.shape[0] == all_num @@ -151,63 +151,63 @@ def test_sample_select_count(): def test_sample_accepts_numpy_scalar_dynamic_batch_size(): - mtp_step = 3 + max_draft_step = 3 probs, b_req_idx = _make_batch_probs( 3, - mtp_step, + max_draft_step, rows=[ [1.0, 0.95, 0.90, 0.10], [1.0, 0.20, 0.80, 0.80], [1.0, 0.99, 0.99, 0.99], ], ) - select = sample_dynamic_mtp_req_mask( + select = sample_dynamic_spec_row_mask( dynamic_batch_size=np.int64(8), b_req_idx=b_req_idx, req_to_next_token_probs=probs, - mtp_step=np.int64(mtp_step), + max_draft_step=np.int64(max_draft_step), ) assert int(select.sum().item()) == 8 def test_sample_topk_by_cumprod_score(): - mtp_step = 3 + max_draft_step = 3 probs, b_req_idx = _make_batch_probs( 3, - mtp_step, + max_draft_step, rows=[ [1.0, 0.95, 0.90, 0.10], [1.0, 0.20, 0.80, 0.80], [1.0, 0.99, 0.99, 0.99], ], ) - flat_probs = _flat_cumprod_probs(b_req_idx, probs, mtp_step) + flat_probs = _flat_cumprod_probs(b_req_idx, probs, max_draft_step) for dynamic_batch_size in [1, 4, 8, 12]: - select = sample_dynamic_mtp_req_mask( + select = sample_dynamic_spec_row_mask( dynamic_batch_size=dynamic_batch_size, b_req_idx=b_req_idx, req_to_next_token_probs=probs.clone(), - mtp_step=mtp_step, + max_draft_step=max_draft_step, ) _assert_topk_mask(select, flat_probs, dynamic_batch_size) def test_sample_picks_highest_cumprod_rows(): - mtp_step = 1 + max_draft_step = 1 probs, b_req_idx = _make_batch_probs( 2, - mtp_step, + max_draft_step, rows=[ [1.0, 0.9], [1.0, 0.1], ], ) - flat_probs = _flat_cumprod_probs(b_req_idx, probs, mtp_step) - select = sample_dynamic_mtp_req_mask( + flat_probs = _flat_cumprod_probs(b_req_idx, probs, max_draft_step) + select = sample_dynamic_spec_row_mask( dynamic_batch_size=2, b_req_idx=b_req_idx, req_to_next_token_probs=probs.clone(), - mtp_step=mtp_step, + max_draft_step=max_draft_step, ) _assert_topk_mask(select, flat_probs, 2) # top-2 scores are both 0.99 at mtp_index==0 (req0 and req1 main rows) @@ -216,14 +216,14 @@ def test_sample_picks_highest_cumprod_rows(): def test_sample_single_request(): - mtp_step = 2 - probs, b_req_idx = _make_batch_probs(1, mtp_step, rows=[[1.0, 0.5, 0.25]]) - flat_probs = _flat_cumprod_probs(b_req_idx, probs, mtp_step) - select = sample_dynamic_mtp_req_mask( + max_draft_step = 2 + probs, b_req_idx = _make_batch_probs(1, max_draft_step, rows=[[1.0, 0.5, 0.25]]) + flat_probs = _flat_cumprod_probs(b_req_idx, probs, max_draft_step) + select = sample_dynamic_spec_row_mask( dynamic_batch_size=2, b_req_idx=b_req_idx, req_to_next_token_probs=probs.clone(), - mtp_step=mtp_step, + max_draft_step=max_draft_step, ) _assert_topk_mask(select, flat_probs, 2) diff --git a/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py b/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py index b28ed9417b..58255bcd36 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py @@ -4,10 +4,10 @@ if not torch.cuda.is_available(): pytest.skip("requires CUDA", allow_module_level=True) -from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_mtp_fa3_decode_params +from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_spec_fa3_decode_params -def _reference_dynamic_mtp_fa3_decode_params(b_req_idx, b_seq_len, b_mark_shared_group, hold_req_id): +def _reference_dynamic_spec_fa3_decode_params(b_req_idx, b_seq_len, b_mark_shared_group, hold_req_id): batch_size = b_req_idx.shape[0] b_req_idx_cpu = b_req_idx.cpu() b_seq_len_cpu = b_seq_len.cpu() @@ -28,7 +28,7 @@ def _reference_dynamic_mtp_fa3_decode_params(b_req_idx, b_seq_len, b_mark_shared @pytest.mark.parametrize("batch_size", [1, 7, 256, 257, 777, 1025, 1537]) -def test_build_dynamic_mtp_fa3_decode_params(batch_size): +def test_build_dynamic_spec_fa3_decode_params(batch_size): hold_req_id = -1 b_req_idx = torch.arange(batch_size, dtype=torch.int32, device="cuda") + 100 b_seq_len = (torch.arange(batch_size, dtype=torch.int32, device="cuda") % 97) + 1 @@ -39,14 +39,14 @@ def test_build_dynamic_mtp_fa3_decode_params(batch_size): if 0 <= pos < batch_size: b_mark_shared_group[pos] = pos % 5 + 1 - actual = build_dynamic_mtp_fa3_decode_params( + actual = build_dynamic_spec_fa3_decode_params( b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mark_shared_group=b_mark_shared_group, att_batch_size=batch_size, hold_req_id=hold_req_id, ) - expected = _reference_dynamic_mtp_fa3_decode_params( + expected = _reference_dynamic_spec_fa3_decode_params( b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mark_shared_group=b_mark_shared_group, @@ -57,21 +57,21 @@ def test_build_dynamic_mtp_fa3_decode_params(batch_size): assert torch.equal(actual_tensor.cpu(), expected_tensor) -def test_build_dynamic_mtp_fa3_decode_params_all_padding(): +def test_build_dynamic_spec_fa3_decode_params_all_padding(): batch_size = 513 hold_req_id = -1 b_req_idx = torch.arange(batch_size, dtype=torch.int32, device="cuda") b_seq_len = torch.arange(batch_size, dtype=torch.int32, device="cuda") + 1 b_mark_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device="cuda") - actual = build_dynamic_mtp_fa3_decode_params( + actual = build_dynamic_spec_fa3_decode_params( b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mark_shared_group=b_mark_shared_group, att_batch_size=batch_size, hold_req_id=hold_req_id, ) - expected = _reference_dynamic_mtp_fa3_decode_params( + expected = _reference_dynamic_spec_fa3_decode_params( b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mark_shared_group=b_mark_shared_group, diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py index 42b2bfab32..036eba833c 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -7,15 +7,17 @@ from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.basemodel.triton_kernel import mtp_utils +from lightllm.utils.envs_utils import get_env_start_args -def test_trim_dynamic_mtp_model_input(monkeypatch): +def test_compact_dynamic_spec_model_input(monkeypatch): monkeypatch.setenv( "LIGHTLLM_START_ARGS", json.dumps( { "diverse_mode": False, "llm_kv_type": "fp16", + "mtp_mode": "eagle3", "mtp_dynamic_verify": True, "mtp_step": 3, "llm_decode_att_backend": "triton", @@ -23,7 +25,7 @@ def test_trim_dynamic_mtp_model_input(monkeypatch): ), ) monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "4") - mtp_utils.get_env_start_args.cache_clear() + get_env_start_args.cache_clear() mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() model_input = ModelInput( @@ -41,6 +43,7 @@ def test_trim_dynamic_mtp_model_input(monkeypatch): mem_indexes=torch.arange(12, dtype=torch.int32, device="cuda") + 100, mem_indexes_cpu=torch.arange(12, dtype=torch.int32, device="cpu") + 100, is_prefill=False, + draft_step=3, multimodal_params=[{"row": i, "images": [], "audios": []} for i in range(12)], mtp_draft_input_hiddens=(torch.arange(12 * 5, dtype=torch.float32, device="cuda").reshape(12, 5) + 0.5), ) @@ -54,11 +57,10 @@ def test_trim_dynamic_mtp_model_input(monkeypatch): device="cuda", ) - trimmed_input, selected_mask = mtp_utils.prepare_dynamic_mtp_model_input( + compacted_input, selected_row_mask = mtp_utils.prepare_dynamic_spec_model_input( model_input=model_input, req_num=3, dynamic_batch_size=8, - req_to_next_token_ids=torch.empty((0,), dtype=torch.int64, device="cuda"), req_to_next_token_probs=req_to_next_token_probs, ) torch.cuda.synchronize() @@ -66,37 +68,39 @@ def test_trim_dynamic_mtp_model_input(monkeypatch): expected_selected_mask = torch.tensor([1, 1, 1, 0, 1, 0, 0, 0, 1, 1, 1, 1], dtype=torch.int32) expected_selected_rows = torch.where(expected_selected_mask == 1)[0] - assert torch.equal(selected_mask.cpu(), expected_selected_mask) - assert trimmed_input.batch_size == 8 - assert trimmed_input.max_q_seq_len == 1 - assert trimmed_input.multimodal_params == [{"images": [], "audios": []}] * 8 + assert torch.equal(selected_row_mask.cpu(), expected_selected_mask) + assert compacted_input.batch_size == 8 + assert compacted_input.max_q_seq_len == 1 + assert compacted_input.multimodal_params == [{"images": [], "audios": []}] * 8 assert torch.equal( - trimmed_input.input_ids.cpu(), torch.arange(12, dtype=torch.int64)[expected_selected_rows] + 1000 + compacted_input.input_ids.cpu(), torch.arange(12, dtype=torch.int64)[expected_selected_rows] + 1000 ) - assert torch.equal(trimmed_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 1, 2, 2, 2, 2], dtype=torch.int32)) - assert torch.equal(trimmed_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 0, 0, 1, 2, 3], dtype=torch.int32)) - assert torch.equal(trimmed_input.b_seq_len.cpu(), torch.tensor([3, 4, 5, 3, 3, 4, 5, 6], dtype=torch.int32)) + assert torch.equal(compacted_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 1, 2, 2, 2, 2], dtype=torch.int32)) + assert torch.equal(compacted_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 0, 0, 1, 2, 3], dtype=torch.int32)) + assert torch.equal(compacted_input.b_seq_len.cpu(), torch.tensor([3, 4, 5, 3, 3, 4, 5, 6], dtype=torch.int32)) assert torch.equal( - trimmed_input.b_position_delta.cpu(), + compacted_input.b_position_delta.cpu(), torch.tensor([200, 201, 202, 204, 208, 209, 210, 211], dtype=torch.int32), ) - assert torch.equal(trimmed_input.b_shared_seq_len.cpu(), torch.tensor([0, 0, 0, 7, 9, 9, 9, 9], dtype=torch.int32)) assert torch.equal( - trimmed_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 1, 0, 0, 0, 4], dtype=torch.int32) + compacted_input.b_shared_seq_len.cpu(), torch.tensor([0, 0, 0, 7, 9, 9, 9, 9], dtype=torch.int32) ) assert torch.equal( - trimmed_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 104, 108, 109, 110, 111], dtype=torch.int32) + compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 1, 0, 0, 0, 4], dtype=torch.int32) + ) + assert torch.equal( + compacted_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 104, 108, 109, 110, 111], dtype=torch.int32) ) # The hot path intentionally keeps the CPU copy unfiltered to avoid a GPU-to-CPU # synchronization. The router frees rejected indexes after its async mask copy. - assert torch.equal(trimmed_input.mem_indexes_cpu, torch.arange(12, dtype=torch.int32) + 100) + assert torch.equal(compacted_input.mem_indexes_cpu, torch.arange(12, dtype=torch.int32) + 100) expected_hiddens = (torch.arange(12 * 5, dtype=torch.float32).reshape(12, 5) + 0.5)[expected_selected_rows] - assert torch.equal(trimmed_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) + assert torch.equal(compacted_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) -def test_trim_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size(monkeypatch): +def test_compaction_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size(monkeypatch): monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "3") mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() @@ -113,21 +117,22 @@ def test_trim_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size(monkey b_mark_shared_group=torch.tensor([0, 0, 0, 0, 5], dtype=torch.int32, device="cuda"), mem_indexes=torch.arange(5, dtype=torch.int32, device="cuda"), mem_indexes_cpu=torch.arange(5, dtype=torch.int32, device="cpu"), + draft_step=4, is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(5)], ) - selected_mask = torch.ones((5,), dtype=torch.int32, device="cuda") + selected_row_mask = torch.ones((5,), dtype=torch.int32, device="cuda") - trimmed_input = mtp_utils._trim_decode_model_input_inplace( + compacted_input = mtp_utils._compact_decode_model_input( model_input=model_input, - selected_mask_gpu=selected_mask, + selected_row_mask=selected_row_mask, dynamic_batch_size=5, ) torch.cuda.synchronize() - assert torch.equal(trimmed_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32)) - assert torch.equal(trimmed_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32)) - assert torch.equal(trimmed_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 0, 2], dtype=torch.int32)) + assert torch.equal(compacted_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32)) + assert torch.equal(compacted_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32)) + assert torch.equal(compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 0, 2], dtype=torch.int32)) def test_mtp_verify_scatter_and_start_locations(): @@ -146,16 +151,16 @@ def test_mtp_verify_scatter_and_start_locations(): device="cuda", ) - mtp_accept_len, accepted_index = mtp_utils.mtp_verify( + spec_accept_len, accepted_index = mtp_utils.mtp_verify( req_to_next_token_ids, b_req_mtp_start_loc, new_next_token_ids, b_req_idx ) mtp_utils.mtp_scatter_next_token_ids( - req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, mtp_accept_len + req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, spec_accept_len ) torch.cuda.synchronize() assert torch.equal(b_req_mtp_start_loc.cpu(), torch.tensor([0, 2], dtype=torch.int64)) - assert torch.equal(mtp_accept_len.cpu(), torch.tensor([1, 1], dtype=torch.int32)) + assert torch.equal(spec_accept_len.cpu(), torch.tensor([1, 1], dtype=torch.int32)) assert torch.equal(accepted_index.cpu(), torch.tensor([1, 0, 1, 0, 0], dtype=torch.int32)) assert torch.equal( req_to_next_token_ids.cpu(), diff --git a/unit_tests/common/speculative/test_config.py b/unit_tests/common/speculative/test_config.py deleted file mode 100644 index cbcd4d6394..0000000000 --- a/unit_tests/common/speculative/test_config.py +++ /dev/null @@ -1,40 +0,0 @@ -from types import SimpleNamespace - -import pytest - -from lightllm.common.speculative.config import SpeculativeConfig - - -@pytest.mark.parametrize( - ("mode", "step", "requested_dynamic", "expected_dynamic", "draft_model_count"), - [ - (None, 0, False, False, 0), - ("vanilla_with_att", 3, False, False, 3), - ("eagle3", 3, True, True, 1), - ("dspark", 0, False, True, 1), - ("dflash", 0, True, False, 1), - ], -) -def test_speculative_config_normalizes_modes( - mode, step, requested_dynamic, expected_dynamic, draft_model_count -): - config = SpeculativeConfig.from_args( - SimpleNamespace(mtp_mode=mode, mtp_step=step, mtp_dynamic_verify=requested_dynamic) - ) - - config.validate() - assert config.dynamic_verify is expected_dynamic - assert config.draft_model_count == draft_model_count - assert config.enabled is (mode is not None) - - -def test_dynamic_eagle3_uses_unit_graph_granularity(): - config = SpeculativeConfig(mode="eagle3", step=7, dynamic_verify=True) - - assert config.get_decode_graph_mtp_step(model_config={}, is_draft_model=False) == 0 - assert config.get_decode_graph_warmup_mtp_step(model_config={}, is_draft_model=False) == 3 - - -def test_invalid_mode_is_rejected(): - with pytest.raises(AssertionError, match="unsupported speculative mode"): - SpeculativeConfig(mode="unknown", step=1).validate() diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py new file mode 100644 index 0000000000..71dfe754c5 --- /dev/null +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -0,0 +1,59 @@ +import torch + +from lightllm.common.basemodel import batch_objs +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.models.qwen3_dspark import model_output as dspark_model_output +from lightllm.models.qwen3_dspark.model import DSparkModelOutputMixin +from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput + + +class _DSparkTestModel(DSparkModelOutputMixin, TpPartBaseModel): + pass + + +def test_dspark_decode_unpad_preserves_output_type_and_slices_dspark_fields(): + model = _DSparkTestModel.__new__(_DSparkTestModel) + output = DSparkModelOutput( + logits=torch.arange(48).view(12, 4), + spec_hidden=torch.arange(36).view(12, 3), + confidence_logits=torch.arange(12).view(3, 4), + draft_token_ids=torch.arange(12), + ) + + unpadded = model._create_unpad_decode_model_output(output, origin_batch_size=8) + + assert isinstance(unpadded, DSparkModelOutput) + assert unpadded.logits.shape == (8, 4) + assert unpadded.spec_hidden.shape == (8, 3) + assert unpadded.confidence_logits.shape == (2, 4) + assert unpadded.draft_token_ids.shape == (8,) + assert output.logits.shape == (12, 4) + assert output.confidence_logits.shape == (3, 4) + assert output.draft_token_ids.shape == (12,) + + +def test_dspark_no_ref_conversion_dispatches_to_dspark_fields(monkeypatch): + monkeypatch.setattr(batch_objs, "tensor_to_no_ref_tensor", torch.clone) + monkeypatch.setattr(dspark_model_output, "tensor_to_no_ref_tensor", torch.clone) + output = DSparkModelOutput( + logits=torch.ones((2, 4)), + spec_hidden=torch.ones((2, 3)), + confidence_logits=torch.ones((1, 2)), + draft_token_ids=torch.ones((2,), dtype=torch.int64), + ) + original_ptrs = ( + output.logits.data_ptr(), + output.spec_hidden.data_ptr(), + output.confidence_logits.data_ptr(), + output.draft_token_ids.data_ptr(), + ) + + output.to_no_ref_tensor() + + converted_ptrs = ( + output.logits.data_ptr(), + output.spec_hidden.data_ptr(), + output.confidence_logits.data_ptr(), + output.draft_token_ids.data_ptr(), + ) + assert all(converted != original for converted, original in zip(converted_ptrs, original_ptrs)) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py new file mode 100644 index 0000000000..c9dd507089 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py @@ -0,0 +1,107 @@ +from types import SimpleNamespace + +import torch + +from lightllm.server.router.model_infer.mode_backend.dp_backend.impl import DPChunkedPrefillBackend +from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal + + +class _RecordingSpecEngine: + def __init__(self): + self.propose_args = None + self.propose_overlap_args = None + self.scatter_args = None + + def propose_next(self, **kwargs): + self.propose_args = kwargs + token_ids = kwargs["next_token_ids"].new_zeros((16, 8)) + return SpecProposal( + token_ids=token_ids, + extra_mem_indexes_cpu=torch.tensor([123], dtype=torch.int32), + ) + + def propose_next_overlap(self, **kwargs): + self.propose_overlap_args = kwargs + row_count = kwargs["real_verify_rows0"] + kwargs["real_verify_rows1"] + token_ids = kwargs["next_token_ids0"].new_zeros((row_count, 8)) + return SpecProposal( + token_ids=token_ids, + extra_mem_indexes_cpu=torch.tensor([456], dtype=torch.int32), + ) + + def scatter_next_tokens(self, **kwargs): + self.scatter_args = kwargs + + +def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.max_draft_step = 7 + backend.spec_engine = _RecordingSpecEngine() + model_input = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + model_output = SimpleNamespace(spec_hidden=torch.randn(16, 4)) + next_token_ids = torch.arange(8, dtype=torch.int64) + real_start_locs = torch.tensor([0], dtype=torch.int32) + real_accept_len = torch.tensor([2], dtype=torch.int32) + + extra_mem = backend._draft_decode_eagle( + model_input=model_input, + model_output=model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=real_start_locs, + spec_accept_len=real_accept_len, + req_num=8, + ) + + propose_args = backend.spec_engine.propose_args + assert propose_args["next_token_ids"].shape == (16,) + assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) + assert torch.equal(propose_args["b_req_mtp_start_loc"], torch.tensor([0, 8], dtype=torch.int32)) + assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) + assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (8, 8) + assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) + + +def test_dp_overlap_eagle_keeps_target_padding_out_of_recurrent_decode(): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.max_draft_step = 7 + backend.spec_engine = _RecordingSpecEngine() + model_input0 = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + model_input1 = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + model_output0 = SimpleNamespace(spec_hidden=torch.randn(16, 4)) + model_output1 = SimpleNamespace(spec_hidden=torch.randn(16, 4)) + next_token_ids = torch.arange(24, dtype=torch.int64) + b_req_idx = torch.arange(24, dtype=torch.int32) + start_locs = torch.tensor([0, 8, 16], dtype=torch.int32) + accept_len = torch.tensor([2, 3, 4], dtype=torch.int32) + + extra_mem = backend._draft_decode_eagle_overlap( + model_input0=model_input0, + model_output0=model_output0, + model_input1=model_input1, + model_output1=model_output1, + b_req_idx=b_req_idx, + next_token_ids=next_token_ids, + spec_accept_len=accept_len, + b_req_mtp_start_loc=start_locs, + req_num0=8, + req_num1=16, + ) + + propose_args = backend.spec_engine.propose_overlap_args + assert propose_args["next_token_ids0"].shape == (16,) + assert propose_args["next_token_ids1"].shape == (16,) + assert propose_args["real_verify_rows0"] == 8 + assert propose_args["real_verify_rows1"] == 16 + assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) + assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) + assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (24, 8) + assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py index e8b3c19719..149a3cccca 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py @@ -4,7 +4,7 @@ if not torch.cuda.is_available(): pytest.skip("requires CUDA", allow_module_level=True) -from lightllm.server.router.model_infer.mode_backend.generic_post_process import _trim_post_sample_tensors +from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import trim_post_sample_tensors def test_trim_post_sample_tensors(): @@ -30,9 +30,9 @@ def test_trim_post_sample_tensors(): out_b_top_ks, out_b_length_penalty_param, out_b_mask_eos_reqs, - ) = _trim_post_sample_tensors( + ) = trim_post_sample_tensors( dynamic_batch_size=dynamic_batch_size, - selected_run_reqs=selected, + selected_row_mask=selected, b_req_idx=b_req_idx, b_temperatures=b_temperatures, b_top_ps=b_top_ps, diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py new file mode 100644 index 0000000000..37a8145e56 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -0,0 +1,62 @@ +from types import SimpleNamespace + +import torch +from lightllm.server.router.model_infer.mode_backend import generic_padded_pre_process, generic_pre_process + + +def test_spec_shared_group_markers_split_requests_and_size_limit(monkeypatch): + monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 2) + b_mtp_index = torch.tensor([0, 1, 2, 0, 1, 0], dtype=torch.int32) + + markers = generic_pre_process.build_spec_shared_group_markers(b_mtp_index=b_mtp_index) + + assert markers.tolist() == [0, 2, 1, 0, 2, 1] + + +def test_spec_shared_group_markers_include_padded_request_rows(monkeypatch): + monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 8) + # The last two groups represent padded requests. Their request ids are both + # HOLD_REQUEST_ID, so mtp_index resets define the actual group boundaries. + b_mtp_index = torch.tensor([0, 1, 2, 0, 1, 2, 0, 1, 2], dtype=torch.int32) + + markers = generic_pre_process.build_spec_shared_group_markers(b_mtp_index=b_mtp_index) + + assert markers.tolist() == [0, 0, 3, 0, 0, 3, 0, 0, 3] + + +def test_padded_decode_builds_spec_metadata_for_real_and_fake_rows(monkeypatch): + max_draft_step = 2 + mem_manager = SimpleNamespace( + HOLD_TOKEN_MEMINDEX=-1, + alloc=lambda size: torch.arange(size, dtype=torch.int32), + ) + req_manager = SimpleNamespace(HOLD_REQUEST_ID=-1, mem_manager=mem_manager) + monkeypatch.setattr( + generic_padded_pre_process, + "g_infer_context", + SimpleNamespace(req_manager=req_manager, radix_cache=None), + ) + monkeypatch.setattr( + generic_padded_pre_process, + "get_env_start_args", + lambda: SimpleNamespace(mtp_step=max_draft_step), + ) + monkeypatch.setattr(generic_padded_pre_process, "enable_diverse_mode_gqa_decode_fast_kernel", lambda: False) + monkeypatch.setattr(generic_padded_pre_process, "enable_dynamic_spec", lambda: True) + monkeypatch.setattr(generic_padded_pre_process, "enable_triton_mtp_kernel", lambda: False) + monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 8) + + req = SimpleNamespace( + req_idx=7, + cur_kv_len=4, + max_draft_step=max_draft_step, + multimodal_params={"images": [], "audios": []}, + get_cur_total_len=lambda: 5, + ) + model_input, _, padded_req_num = generic_padded_pre_process.padded_prepare_decode_inputs( + req_objs=[req], dest_batch_size=3 + ) + + assert padded_req_num == 2 + assert model_input.b_mtp_index.tolist() == [0, 1, 2, 0, 1, 2, 0, 1, 2] + assert model_input.b_mark_shared_group.tolist() == [0, 0, 3, 0, 0, 3, 0, 0, 3] diff --git a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py new file mode 100644 index 0000000000..57501d3a19 --- /dev/null +++ b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py @@ -0,0 +1,88 @@ +from types import SimpleNamespace + +import torch + +from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer + + +class _DraftModel: + def __init__(self): + self.extend_batch_sizes = None + self.decode_batch_sizes = [] + + def microbatch_overlap_prefill(self, input0, input1): + self.extend_batch_sizes = (input0.batch_size, input1.batch_size) + return tuple( + ModelOutput( + logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), + spec_hidden=torch.ones((model_input.batch_size, 2)), + ) + for model_input in (input0, input1) + ) + + def microbatch_overlap_decode(self, input0, input1): + self.decode_batch_sizes.append((input0.batch_size, input1.batch_size)) + return tuple( + ModelOutput( + logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), + spec_hidden=torch.ones((model_input.batch_size, 2)), + ) + for model_input in (input0, input1) + ) + + +def _target_input(batch_size): + return SimpleNamespace( + batch_size=batch_size, + b_seq_len=torch.arange(batch_size, dtype=torch.int32) + 4, + b_req_idx=torch.arange(batch_size, dtype=torch.int32), + b_position_delta=torch.zeros(batch_size, dtype=torch.int32), + max_kv_seq_len=16, + max_cache_len=16, + ) + + +def test_overlap_eagle_extends_verify_rows_then_decodes_logical_batch(monkeypatch): + # Keep this topology test CPU-only; production allocates the same tensor + # through the shared CUDA KV manager. + monkeypatch.setattr(torch.Tensor, "cuda", lambda self, **_: self) + draft_model = _DraftModel() + backend = SimpleNamespace( + max_draft_step=2, + draft_models=[draft_model], + model=SimpleNamespace( + req_manager=SimpleNamespace( + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + ) + ), + _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), + ) + engine = SimpleNamespace( + backend=backend, + enable_dynamic_spec=False, + alloc_extra_mem_indexes=lambda token_count: torch.arange(token_count, dtype=torch.int32), + ) + proposer = EagleMTPProposer(engine) + model_input0 = _target_input(batch_size=6) + model_input1 = _target_input(batch_size=6) + + proposal = proposer.propose_next_overlap( + main_model_input0=model_input0, + main_model_output0=ModelOutput(logits=torch.empty((6, 1)), spec_hidden=torch.ones((6, 2))), + next_token_ids0=torch.arange(6, dtype=torch.int64), + real_verify_rows0=3, + accept_len0=torch.tensor([2, 1], dtype=torch.int32), + main_model_input1=model_input1, + main_model_output1=ModelOutput(logits=torch.empty((6, 1)), spec_hidden=torch.ones((6, 2))), + next_token_ids1=torch.arange(10, 16, dtype=torch.int64), + real_verify_rows1=6, + accept_len1=torch.tensor([1, 3], dtype=torch.int32), + draft_step=2, + ) + + assert draft_model.extend_batch_sizes == (6, 6) + assert draft_model.decode_batch_sizes == [(2, 2)] + assert proposal.token_ids.shape == (9, 3) + assert torch.equal(proposal.token_ids[:, 0], torch.tensor([0, 1, 2, 10, 11, 12, 13, 14, 15])) + assert proposal.extra_mem_indexes_cpu.shape == (3,) diff --git a/unit_tests/server/router/model_infer/speculative/test_planner.py b/unit_tests/server/router/model_infer/speculative/test_planner.py index b5987e1396..58f0bb9ab2 100644 --- a/unit_tests/server/router/model_infer/speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/speculative/test_planner.py @@ -3,30 +3,26 @@ import numpy as np from lightllm.server.router.model_infer.speculative.planner import ( - DSparkDynamicMTPPlanner, - DynamicMTPPlanner, - FixedMTPPlanner, + DSparkDynamicSpecPlanner, + DynamicSpecPlanner, + FixedSpecPlanner, ) def test_fixed_planner_returns_static_plan(): - plan = FixedMTPPlanner(mtp_step=3).plan(req_num=4, original_batch_size=16) + plan = FixedSpecPlanner(max_draft_step=3).plan(req_num=4, original_batch_size=16) assert not plan.is_dynamic assert plan.dynamic_batch_size is None assert plan.draft_step == plan.pre_draft_step == 3 - assert plan.selection_mode == "none" assert not plan.skip_verify_sync def test_dynamic_planner_stays_full_width_until_costs_are_profiled(): - plan = DynamicMTPPlanner(mtp_step=3, use_random_mode=False).plan( - req_num=2, original_batch_size=8 - ) + plan = DynamicSpecPlanner(max_draft_step=3, use_random_mode=False).plan(req_num=2, original_batch_size=8) assert plan.dynamic_batch_size == 8 assert plan.draft_step == plan.pre_draft_step == 3 - assert plan.selection_mode == "observe" def test_planner_does_not_reset_process_global_random_state(): @@ -34,13 +30,13 @@ def test_planner_does_not_reset_process_global_random_state(): expected = random.random() random.seed(2027) - DynamicMTPPlanner(mtp_step=3) + DynamicSpecPlanner(max_draft_step=3) assert random.random() == expected def test_dspark_applies_confidence_capacity_after_two_step_delay(): - planner = DSparkDynamicMTPPlanner(mtp_step=3) + planner = DSparkDynamicSpecPlanner(max_draft_step=3) planner.update_infer_cost(batch_size=2, infer_cost_ms=1.0, is_draft_model=False) planner.update_infer_cost(batch_size=4, infer_cost_ms=1.1, is_draft_model=False) planner.update_infer_cost(batch_size=8, infer_cost_ms=10.0, is_draft_model=False) @@ -57,7 +53,7 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): def test_topk_prefix_sums_only_computes_requested_counts(): - result = DSparkDynamicMTPPlanner._topk_prefix_sums( + result = DSparkDynamicSpecPlanner._topk_prefix_sums( values=np.asarray([0.1, 0.9, 0.4, 0.7]), counts=[0, 2, 4], ) diff --git a/unit_tests/server/test_api_start_spec_config.py b/unit_tests/server/test_api_start_spec_config.py deleted file mode 100644 index 8982969703..0000000000 --- a/unit_tests/server/test_api_start_spec_config.py +++ /dev/null @@ -1,42 +0,0 @@ -from types import SimpleNamespace - -from lightllm.common.speculative.config import SpeculativeConfig -from lightllm.server import api_start - - -def test_block_mode_derives_mtp_step_from_checkpoint(monkeypatch): - draft_config = { - "architectures": ["Qwen3DSparkModel"], - "block_size": 4, - "target_layer_ids": [8, 16, 24], - "mask_token_id": 151665, - "markov_rank": 0, - "enable_confidence_head": True, - "confidence_head_with_markov": False, - } - monkeypatch.setattr( - api_start.PretrainedConfig, - "get_config_dict", - lambda _model_dir: (draft_config, {}), - ) - args = SimpleNamespace(mtp_draft_model_dir=["/models/dspark"], mtp_step=0) - config = SpeculativeConfig(mode="dspark", step=0, dynamic_verify=True) - - normalized = api_start.normalize_block_mtp_step_from_first_draft_config(args, config) - - assert args.mtp_step == 4 - assert normalized.step == 4 - normalized.validate() - - -def test_non_block_mode_keeps_explicit_step(monkeypatch): - monkeypatch.setattr( - api_start.PretrainedConfig, - "get_config_dict", - lambda _model_dir: (_ for _ in ()).throw(AssertionError("must not load config")), - ) - args = SimpleNamespace(mtp_draft_model_dir=["/models/eagle3"], mtp_step=3) - config = SpeculativeConfig(mode="eagle3", step=3, dynamic_verify=True) - - assert api_start.normalize_block_mtp_step_from_first_draft_config(args, config) is config - assert args.mtp_step == 3 diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py new file mode 100644 index 0000000000..67607433d1 --- /dev/null +++ b/unit_tests/utils/test_speculative_utils.py @@ -0,0 +1,208 @@ +import json +from importlib import import_module +from types import SimpleNamespace + +import pytest + +import lightllm.common.basemodel.attention.base_att as base_att_module +from lightllm.common.basemodel.attention.base_att import BaseAttBackend +from lightllm.models import get_draft_model_class +from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight +from lightllm.utils import envs_utils + + +@pytest.mark.parametrize( + "module_name, class_name", + [ + ("lightllm.models.deepseek_mtp.model", "Deepseek3MTPModel"), + ("lightllm.models.glm4_moe_lite_mtp.model", "Glm4MoeLiteMTPModel"), + ("lightllm.models.mistral_mtp.model", "MistralMTPModel"), + ("lightllm.models.qwen3_moe_mtp.model", "Qwen3MOEMTPModel"), + ("lightllm.models.qwen3_5_mtp.model", "Qwen3_5MTPModel"), + ("lightllm.models.qwen3_5_moe_mtp.model", "Qwen3_5MoeMTPModel"), + ("lightllm.models.qwen3_eagle.model", "Qwen3EagleModel"), + ("lightllm.models.qwen3_dflash.model", "Qwen3DFlashModel"), + ("lightllm.models.qwen3_5_dflash.model", "Qwen3_5DFlashModel"), + ("lightllm.models.qwen3_dspark.model", "Qwen3DSparkModel"), + ("lightllm.models.qwen3_5_dspark.model", "Qwen3_5DSparkModel"), + ], +) +def test_spec_draft_model_class_is_marked(module_name, class_name): + model_class = getattr(import_module(module_name), class_name) + assert model_class.is_mtp_draft_model is True + + +def test_qwen3_eagle_uses_layers_checkpoint_prefix(): + layer_weight = Qwen3EagleTransformerLayerWeight.__new__(Qwen3EagleTransformerLayerWeight) + layer_weight.layer_num_ = 2 + + layer_weight._init_weight_names() + + assert layer_weight._q_weight_name == "layers.2.self_attn.q_proj.weight" + assert layer_weight._hidden_norm_weight_name == "layers.2.hidden_norm.weight" + + +@pytest.mark.parametrize( + "mtp_mode, is_draft_model, draft_step, dynamic_spec, expected", + [ + ("dspark", False, 7, True, True), + ("dspark", True, 7, True, False), + ("vanilla_with_att", True, 7, True, True), + ("eagle3", True, 0, True, False), + ("eagle3", False, 7, False, False), + ], +) +def test_attention_backend_selects_dynamic_spec_layout( + monkeypatch, + mtp_mode, + is_draft_model, + draft_step, + dynamic_spec, + expected, +): + monkeypatch.setattr(base_att_module, "enable_dynamic_spec", lambda: dynamic_spec) + monkeypatch.setattr(base_att_module, "get_env_start_args", lambda: SimpleNamespace(mtp_mode=mtp_mode)) + backend = SimpleNamespace(model=SimpleNamespace(is_mtp_draft_model=is_draft_model)) + infer_state = SimpleNamespace(draft_step=draft_step) + + assert BaseAttBackend.uses_dynamic_spec_verify_layout(backend, infer_state) is expected + + +@pytest.mark.parametrize( + "spec_mode, architecture, is_linear_att_mixed_model, expected_class_name", + [ + ("dflash", "Qwen3DFlashModel", False, "Qwen3DFlashModel"), + ("dflash", "Qwen3DSparkModel", False, "Qwen3DFlashModel"), + ("dflash", "Qwen3_5DFlashModel", True, "Qwen3_5DFlashModel"), + ("dspark", "Qwen3DSparkModel", False, "Qwen3DSparkModel"), + ("dspark", "Qwen3DSparkModel", True, "Qwen3_5DSparkModel"), + ("eagle3", "Qwen3Eagle3Model", False, "Qwen3EagleModel"), + ], +) +def test_draft_model_registry( + spec_mode, + architecture, + is_linear_att_mixed_model, + expected_class_name, +): + model_class = get_draft_model_class( + model_cfg={"model_type": "qwen3", "architectures": [architecture]}, + spec_mode=spec_mode, + is_linear_att_mixed_model=is_linear_att_mixed_model, + ) + + assert model_class.__name__ == expected_class_name + + +def test_dspark_env_helper_enables_dynamic_spec_without_legacy_flag(monkeypatch): + monkeypatch.setenv( + "LIGHTLLM_START_ARGS", + json.dumps( + { + "mtp_mode": "dspark", + "mtp_dynamic_verify": False, + "mtp_step": 4, + } + ), + ) + envs_utils.get_env_start_args.cache_clear() + envs_utils.enable_dynamic_spec.cache_clear() + + assert envs_utils.enable_dynamic_spec() + + +@pytest.mark.parametrize( + "mtp_mode, mtp_dynamic_verify, expected", + [ + ("dspark", False, True), + ("dflash", True, False), + ("eagle3", True, True), + ("eagle3", False, False), + ], +) +def test_env_helper_uses_normalized_dynamic_spec(monkeypatch, mtp_mode, mtp_dynamic_verify, expected): + monkeypatch.setenv( + "LIGHTLLM_START_ARGS", + json.dumps( + { + "mtp_mode": mtp_mode, + "mtp_step": 7, + "mtp_dynamic_verify": mtp_dynamic_verify, + } + ), + ) + envs_utils.get_env_start_args.cache_clear() + envs_utils.enable_dynamic_spec.cache_clear() + + assert envs_utils.enable_dynamic_spec() is expected + + +def test_draft_model_registry_rejects_unsupported_architecture(): + with pytest.raises(ValueError, match="Unsupported speculative draft model"): + get_draft_model_class( + model_cfg={"model_type": "gemma4", "architectures": ["Gemma4DSparkModel"]}, + spec_mode="dspark", + ) + + +def test_draft_model_registry_validates_target_family(): + with pytest.raises(ValueError, match="linear-attention mixed targets"): + get_draft_model_class( + model_cfg={"model_type": "qwen3", "architectures": ["Qwen3DFlashModel"]}, + spec_mode="dflash", + is_linear_att_mixed_model=True, + ) + + +def test_dflash_added_kv_layers_come_from_draft_config(tmp_path): + config_path = tmp_path / "config.json" + config_path.write_text(json.dumps({"num_hidden_layers": 5})) + + envs_utils.get_env_start_args.cache_clear() + envs_utils.get_added_mtp_kv_layer_num.cache_clear() + envs_utils.set_env_start_args( + { + "mtp_mode": "dflash", + "mtp_step": 7, + "mtp_dynamic_verify": False, + "mtp_draft_model_dir": [str(tmp_path)], + } + ) + + assert envs_utils.get_added_mtp_kv_layer_num() == 5 + + +def test_dspark_added_kv_layers_come_from_draft_config(tmp_path): + config_path = tmp_path / "config.json" + config_path.write_text(json.dumps({"num_hidden_layers": 6})) + + envs_utils.get_env_start_args.cache_clear() + envs_utils.get_added_mtp_kv_layer_num.cache_clear() + envs_utils.set_env_start_args( + { + "mtp_mode": "dspark", + "mtp_step": 7, + "mtp_dynamic_verify": True, + "mtp_draft_model_dir": [str(tmp_path)], + } + ) + + assert envs_utils.get_added_mtp_kv_layer_num() == 6 + + +def test_eagle3_added_kv_layers_come_from_draft_config(tmp_path): + config_path = tmp_path / "config.json" + config_path.write_text(json.dumps({"num_hidden_layers": 2})) + + envs_utils.get_env_start_args.cache_clear() + envs_utils.get_added_mtp_kv_layer_num.cache_clear() + envs_utils.set_env_start_args( + { + "mtp_mode": "eagle3", + "mtp_step": 7, + "mtp_dynamic_verify": False, + "mtp_draft_model_dir": [str(tmp_path)], + } + ) + + assert envs_utils.get_added_mtp_kv_layer_num() == 2 From 6aea8ce632ca0a9afe74ef875044288a0bc4a312 Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Mon, 10 Aug 2026 09:08:31 +0000 Subject: [PATCH 005/103] fix: read Qwen3.5 target layers from text config --- .../model_infer/mode_backend/base_backend.py | 5 +++-- unit_tests/utils/test_speculative_utils.py | 16 ++++++++++++++++ 2 files changed, 19 insertions(+), 2 deletions(-) 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 6841eff919..837cf177fe 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -473,6 +473,9 @@ def _target_hidden_layer_ids(self, model_cfg: dict): if self.args.mtp_mode not in ("eagle3", "dspark", "dflash"): return None + model_cfg = model_cfg.get("text_config", model_cfg) + layer_num = int(model_cfg.get("num_hidden_layers", model_cfg.get("n_layer"))) + draft_model_dirs = self.args.mtp_draft_model_dir assert draft_model_dirs draft_cfg, _ = PretrainedConfig.get_config_dict(draft_model_dirs[0]) @@ -480,10 +483,8 @@ def _target_hidden_layer_ids(self, model_cfg: dict): if target_layer_ids is None and self.args.mtp_mode == "dflash": target_layer_ids = draft_cfg.get("dflash_config", {}).get("target_layer_ids") if target_layer_ids is None: - layer_num = int(model_cfg.get("num_hidden_layers", model_cfg.get("n_layer"))) target_layer_ids = [1, layer_num // 2 - 1, layer_num - 4] - layer_num = int(model_cfg.get("num_hidden_layers", model_cfg.get("n_layer"))) target_layer_ids = tuple(int(layer_id) for layer_id in target_layer_ids) assert target_layer_ids and all( 0 <= layer_id < layer_num for layer_id in target_layer_ids diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 67607433d1..3ee3f36e7f 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -5,9 +5,11 @@ import pytest import lightllm.common.basemodel.attention.base_att as base_att_module +import lightllm.server.router.model_infer.mode_backend.base_backend as base_backend_module from lightllm.common.basemodel.attention.base_att import BaseAttBackend from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight +from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.utils import envs_utils @@ -154,6 +156,20 @@ def test_draft_model_registry_validates_target_family(): ) +def test_target_hidden_layer_ids_reads_text_config(monkeypatch): + monkeypatch.setattr( + base_backend_module.PretrainedConfig, + "get_config_dict", + lambda _: ({"target_layer_ids": [1, 20, 36]}, {}), + ) + backend = ModeBackend.__new__(ModeBackend) + backend.args = SimpleNamespace(mtp_mode="dspark", mtp_draft_model_dir=["/models/dspark"]) + + layer_ids = backend._target_hidden_layer_ids({"model_type": "qwen3_5", "text_config": {"num_hidden_layers": 40}}) + + assert layer_ids == (1, 20, 36) + + def test_dflash_added_kv_layers_come_from_draft_config(tmp_path): config_path = tmp_path / "config.json" config_path.write_text(json.dumps({"num_hidden_layers": 5})) From 3f2ee165489932317f879c8ecde890f6243a8228 Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Tue, 11 Aug 2026 09:07:38 +0000 Subject: [PATCH 006/103] refactor: simplify LightSpec draft integration --- .../common/basemodel/attention/base_att.py | 7 +- lightllm/common/basemodel/attention/fa3/fp.py | 19 +- .../common/basemodel/attention/fa3/mla.py | 19 +- .../common/basemodel/attention/linear/gdn.py | 10 +- lightllm/common/basemodel/basemodel.py | 20 +- lightllm/common/basemodel/hidden_collector.py | 49 ++-- .../meta_weights/embedding_weight.py | 13 +- .../linear_att_cache_manager/config_objs.py | 4 +- lightllm/common/req_manager.py | 24 +- lightllm/models/__init__.py | 55 +--- lightllm/models/deepseek_mtp/model.py | 5 + lightllm/models/glm4_moe_lite_mtp/model.py | 5 + lightllm/models/mistral_mtp/model.py | 5 + lightllm/models/qwen3_5_dflash/model.py | 63 ++--- lightllm/models/qwen3_5_dspark/model.py | 55 ++-- lightllm/models/qwen3_5_moe_mtp/model.py | 5 + lightllm/models/qwen3_5_mtp/model.py | 5 + lightllm/models/qwen3_dflash/model.py | 26 +- lightllm/models/qwen3_dspark/infer_struct.py | 8 + .../layer_infer/post_layer_infer.py | 261 ++++-------------- .../pre_and_post_layer_weight.py | 54 ++-- lightllm/models/qwen3_dspark/model.py | 52 ++-- lightllm/models/qwen3_eagle/model.py | 2 + lightllm/models/qwen3_moe_mtp/model.py | 5 + lightllm/models/registry.py | 47 +++- lightllm/server/api_cli.py | 2 - lightllm/server/api_start.py | 31 +-- lightllm/server/core/objs/start_args_type.py | 2 - .../server/router/model_infer/infer_batch.py | 12 +- .../model_infer/mode_backend/base_backend.py | 75 +---- .../mode_backend/chunked_prefill/impl.py | 16 +- .../mode_backend/diverse_backend/impl.py | 14 +- .../mode_backend/dp_backend/impl.py | 17 +- .../generic_padded_pre_process.py | 13 +- .../mode_backend/generic_post_process.py | 18 +- .../mode_backend/generic_pre_process.py | 5 +- .../router/model_infer/speculative/engine.py | 5 +- .../speculative/proposers/__init__.py | 2 +- .../speculative/proposers/dflash.py | 27 +- .../speculative/proposers/eagle3.py | 11 +- .../speculative/proposers/eagle_mtp.py | 22 +- lightllm/utils/envs_utils.py | 43 +-- .../models/test_qwen3_dspark_model_output.py | 73 ++++- .../mode_backend/test_generic_pre_process.py | 5 +- .../speculative/test_eagle_overlap.py | 17 ++ unit_tests/utils/test_speculative_utils.py | 205 +++++++++----- 46 files changed, 708 insertions(+), 725 deletions(-) create mode 100644 lightllm/models/qwen3_dspark/infer_struct.py diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index 063cd3ccf1..da897faef0 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -3,7 +3,7 @@ from dataclasses import dataclass from typing import Optional, TYPE_CHECKING, Tuple, Union, Dict -from lightllm.utils.envs_utils import enable_dynamic_spec, get_env_start_args +from lightllm.utils.envs_utils import get_env_start_args if TYPE_CHECKING: from lightllm.common.basemodel.basemodel import TpPartBaseModel @@ -40,12 +40,13 @@ def create_att_decode_state(self) -> "BaseDecodeAttState": raise NotImplementedError("not impl") def uses_dynamic_spec_verify_layout(self, infer_state: "InferStateInfo") -> bool: - if infer_state.draft_step == 0 or not enable_dynamic_spec(): + args = get_env_start_args() + if infer_state.draft_step == 0 or not args.mtp_dynamic_verify: return False # Target verification may compact each request to a different row count. # Block draft forwards still use their checkpoint-defined fixed layout. - return get_env_start_args().mtp_mode not in ("dspark", "dflash") or not self.model.is_mtp_draft_model + return args.mtp_mode not in ("dspark", "dflash") or not self.model.is_mtp_draft_model def _find_layer_index( self, k: torch.Tensor, v: torch.Tensor, att_state: Union["BasePrefillAttState", "BaseDecodeAttState"] diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index 89fdc61389..7379957b7e 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -9,6 +9,7 @@ page_table_copy, ) from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor +from lightllm.utils.envs_utils import get_env_start_args class Fa3AttBackend(BaseAttBackend): @@ -22,13 +23,19 @@ def get_page_table_buffer(self): """ model = self.model if not hasattr(self, "_shared_page_table_buffer"): + max_att_batch_size = model.graph_max_batch_size + if not get_env_start_args().mtp_dynamic_verify: + # FA3 merges each fixed speculative block into one attention sequence. + max_att_batch_size //= model.decode_batch_multiplier + + buffer_count = 2 if model.args.enable_decode_microbatch_overlap else 1 self._shared_page_table_buffer = [ - torch.empty(model.graph_max_batch_size * model.graph_max_len_in_batch, dtype=torch.int32).to( - get_current_device_id() - ), - torch.empty(model.graph_max_batch_size * model.graph_max_len_in_batch, dtype=torch.int32).to( - get_current_device_id() - ), + torch.empty( + max_att_batch_size * model.graph_max_len_in_batch, + dtype=torch.int32, + device=get_current_device_id(), + ) + for _ in range(buffer_count) ] return self._shared_page_table_buffer diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index c07cea8d4c..2e7f127fbb 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -7,6 +7,7 @@ from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_spec_fa3_decode_params, page_table_copy from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor from lightllm.utils.sgl_utils import flash_attn_varlen_func +from lightllm.utils.envs_utils import get_env_start_args class MlaFa3AttBackend(BaseAttBackend): @@ -20,13 +21,19 @@ def get_page_table_buffer(self): """ model = self.model if not hasattr(self, "_shared_page_table_buffer"): + max_att_batch_size = model.graph_max_batch_size + if not get_env_start_args().mtp_dynamic_verify: + # FA3 merges each fixed speculative block into one attention sequence. + max_att_batch_size //= model.decode_batch_multiplier + + buffer_count = 2 if model.args.enable_decode_microbatch_overlap else 1 self._shared_page_table_buffer = [ - torch.empty(model.graph_max_batch_size * model.graph_max_len_in_batch, dtype=torch.int32).to( - get_current_device_id() - ), - torch.empty(model.graph_max_batch_size * model.graph_max_len_in_batch, dtype=torch.int32).to( - get_current_device_id() - ), + torch.empty( + max_att_batch_size * model.graph_max_len_in_batch, + dtype=torch.int32, + device=get_current_device_id(), + ) + for _ in range(buffer_count) ] return self._shared_page_table_buffer diff --git a/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index 2441e1e52c..918919ba32 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -29,7 +29,7 @@ def __init__(self, model: "TpPartBaseModel"): def _init_linear_layer_metadata(self, network_config, tp_world_size): - self.max_draft_step = get_env_start_args().mtp_step + self.mtp_step = get_env_start_args().mtp_step # Linear attention specific dimensions self.num_v_heads = network_config["linear_num_value_heads"] @@ -112,12 +112,12 @@ class LinearAttPrefillAttState(BasePrefillAttState): def init_state(self): backend: LinearAttBackend = self.backend - max_draft_step = backend.max_draft_step + mtp_step = backend.mtp_step # 每次 _prefill 都会在 runtime infer_state 上调用 init_state。 # prefill cuda graph 回调必须走 new_infer_state.prefill_att_state1, # 才能读到这里按当前 batch(含 token padding 后的 dummy request)更新的索引。 self.b_conv_buffer_idx = self.infer_state.b_req_idx - self.b_ssm_buffer_idx = self.infer_state.b_req_idx * (max_draft_step + 1) + self.b_ssm_buffer_idx = self.infer_state.b_req_idx * (mtp_step + 1) return def prefill_att( @@ -140,8 +140,8 @@ def prefill_att( conv_states, ssm_states = self.infer_state.req_manager.get_mamba_cache(layer_num) # 在开启了mtp的时候,conv 状态的最后一维可能存在冗余的部分,需要进行切片对齐。 # prefill 模式下,使用不到这几个维度,所以需要扣除掉, - if backend.max_draft_step > 0: - conv_states = conv_states[:, :, : -backend.max_draft_step] + if backend.mtp_step > 0: + conv_states = conv_states[:, :, : -backend.mtp_step] mixed_qkv, z, b, a = backend._split_qkvzba(mixed_qkvzba) core_attn_out = self._gdn_prefill_kernel( mixed_qkv, conv_states, ssm_states, a, b, self.infer_state, layer_weight diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 85ca25cfd6..5ecd322351 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -37,7 +37,6 @@ set_model_init_status, enable_diverse_mode_gqa_decode_fast_kernel, enable_full_att_decode_tune, - enable_dynamic_spec, enable_triton_mtp_kernel, ) from lightllm.common.triton_utils.autotuner import Autotuner @@ -137,7 +136,7 @@ def __init__(self, kvargs): logger.info(f"use prefill att backend1: {self.prefill_att_backend1.__class__.__name__}") logger.info(f"use decode att backend1: {self.decode_att_backend1.__class__.__name__}") - self._init_hidden_collector(hidden_layer_ids=kvargs.get("hidden_layer_ids")) + self._init_hidden_collector() self._autotune_warmup() self._full_att_decode_autotune() self._init_padded_req() @@ -274,7 +273,9 @@ def _init_att_backend1(self): return def _init_cudagraph(self): - batch_multiplier = 1 if enable_dynamic_spec() and not self.is_mtp_draft_model else self.decode_batch_multiplier + batch_multiplier = ( + 1 if self.args.mtp_dynamic_verify and not self.is_mtp_draft_model else self.decode_batch_multiplier + ) self.graph = ( None if self.disable_cudagraph @@ -283,7 +284,7 @@ def _init_cudagraph(self): max_len_in_batch=self.graph_max_len_in_batch, tp_world_size=self.tp_world_size_, batch_multiplier=batch_multiplier, - capture_infer_cost=enable_dynamic_spec(), + capture_infer_cost=self.args.mtp_dynamic_verify, ) ) if self.graph is not None: @@ -343,7 +344,9 @@ def _full_att_decode_autotune(self): from lightllm.utils.sgl_utils import fa3_decode_autotune - batch_multiplier = 1 if enable_dynamic_spec() and not self.is_mtp_draft_model else self.decode_batch_multiplier + batch_multiplier = ( + 1 if self.args.mtp_dynamic_verify and not self.is_mtp_draft_model else self.decode_batch_multiplier + ) cuda_graph_batch_sizes = CudaGraph.gen_cuda_graph_batch_sizes( max_batch_size=self.graph_max_batch_size, tp_world_size=self.tp_world_size_, @@ -355,14 +358,13 @@ def _full_att_decode_autotune(self): def _init_custom(self): pass - def _init_hidden_collector(self, hidden_layer_ids): + def _init_hidden_collector(self): microbatch_count = ( 2 if self.args.enable_prefill_microbatch_overlap or self.args.enable_decode_microbatch_overlap else 1 ) self.hidden_collector = HiddenCollector( model=self, spec_mode=self.args.mtp_mode, - layer_ids=hidden_layer_ids, microbatch_count=microbatch_count, ) @@ -402,7 +404,7 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) if enable_diverse_mode_gqa_decode_fast_kernel(): infer_state.b_shared_seq_len = model_input.b_shared_seq_len infer_state.b_mark_shared_group = model_input.b_mark_shared_group - elif enable_dynamic_spec() or enable_triton_mtp_kernel(): + elif self.args.mtp_dynamic_verify or enable_triton_mtp_kernel(): infer_state.b_mark_shared_group = model_input.b_mark_shared_group infer_state.multimodal_params = model_input.multimodal_params @@ -475,7 +477,7 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.b_mark_shared_group = F.pad( new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=1 ) - elif enable_dynamic_spec() or enable_triton_mtp_kernel(): + elif self.args.mtp_dynamic_verify or enable_triton_mtp_kernel(): assert new_model_input.b_mark_shared_group is not None new_model_input.b_mark_shared_group = F.pad( new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=0 diff --git a/lightllm/common/basemodel/hidden_collector.py b/lightllm/common/basemodel/hidden_collector.py index a03fa90e59..64d67fe69f 100644 --- a/lightllm/common/basemodel/hidden_collector.py +++ b/lightllm/common/basemodel/hidden_collector.py @@ -3,6 +3,9 @@ from typing import Iterable, List, Optional import torch +from transformers.configuration_utils import PretrainedConfig + +from lightllm.utils.envs_utils import get_env_start_args def unpad_collected_hidden(hidden: Optional[torch.Tensor], token_count: int) -> Optional[torch.Tensor]: @@ -42,13 +45,29 @@ def finish( class LayerHiddenCollector(NoopHiddenCollector): """Collects selected decoder-layer outputs for an intermediate-hidden draft.""" - def __init__(self, model, layer_ids: Iterable[int]) -> None: + def __init__(self, model, layer_ids: Optional[Iterable[int]] = None) -> None: self.model = model self.layer_num = model.layers_num - self.layer_ids = frozenset(int(layer_id) for layer_id in layer_ids) - assert self.layer_ids, "layer hidden collector requires at least one layer id" + self.layer_ids = self._resolve_layer_ids(layer_ids) self.layer_hiddens: List[torch.Tensor] = [] + def _resolve_layer_ids(self, layer_ids: Optional[Iterable[int]]) -> frozenset[int]: + if layer_ids is None: + draft_model_dirs = get_env_start_args().mtp_draft_model_dir + assert draft_model_dirs + draft_config, _ = PretrainedConfig.get_config_dict(draft_model_dirs[0]) + layer_ids = draft_config.get("target_layer_ids") + if layer_ids is None: + layer_ids = draft_config.get("dflash_config", {}).get("target_layer_ids") + if layer_ids is None: + layer_ids = [1, self.layer_num // 2 - 1, self.layer_num - 4] + + resolved_layer_ids = frozenset(int(layer_id) for layer_id in layer_ids) + assert resolved_layer_ids and all( + 0 <= layer_id < self.layer_num for layer_id in resolved_layer_ids + ), f"invalid target_layer_ids={resolved_layer_ids} for target layer_num={self.layer_num}" + return resolved_layer_ids + def add(self, layer_index: int, hidden: torch.Tensor) -> None: if layer_index not in self.layer_ids: return @@ -92,21 +111,19 @@ def __init__( microbatch_count: int = 1, ) -> None: assert microbatch_count > 0 - layer_ids = None if layer_ids is None else tuple(layer_ids) - - collector_type = NoopHiddenCollector - collector_kwargs = {} if spec_mode is not None: assert model is not None - if model.is_mtp_draft_model: - if spec_mode not in ("dspark", "dflash"): - collector_type = FinalHiddenCollector - elif spec_mode in ("eagle3", "dspark", "dflash"): - assert layer_ids is not None - collector_type = LayerHiddenCollector - collector_kwargs = {"model": model, "layer_ids": layer_ids} - else: - collector_type = FinalHiddenCollector + + collector_kwargs = {} + if spec_mode is None: + collector_type = NoopHiddenCollector + elif model.is_mtp_draft_model: + collector_type = NoopHiddenCollector if spec_mode in ("dspark", "dflash") else FinalHiddenCollector + elif spec_mode not in ("eagle3", "dspark", "dflash"): + collector_type = FinalHiddenCollector + else: + collector_type = LayerHiddenCollector + collector_kwargs = {"model": model, "layer_ids": layer_ids} self.collectors = tuple(collector_type(**collector_kwargs) for _ in range(microbatch_count)) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/embedding_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/embedding_weight.py index d94a4c709b..32267d90a5 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/embedding_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/embedding_weight.py @@ -8,8 +8,16 @@ class EmbeddingWeight(BaseWeightTpl, PlatformAwareOp): - def __init__(self, dim: int, vocab_size: int, weight_name: str, data_type: torch.dtype): - super().__init__() + def __init__( + self, + dim: int, + vocab_size: int, + weight_name: str, + data_type: torch.dtype, + tp_rank: Optional[int] = None, + tp_world_size: Optional[int] = None, + ): + super().__init__(tp_rank=tp_rank, tp_world_size=tp_world_size, data_type=data_type) self.dim = dim self.vocab_size = vocab_size # 计算 split_indexes @@ -17,7 +25,6 @@ def __init__(self, dim: int, vocab_size: int, weight_name: str, data_type: torch self.tp_vocab_start_id = int(split_indexes[self.tp_rank_]) self.tp_vocab_end_id = int(split_indexes[self.tp_rank_ + 1]) self.weight_name: str = weight_name - self.data_type_ = data_type self._create_weight() def _create_weight(self): diff --git a/lightllm/common/linear_att_cache_manager/config_objs.py b/lightllm/common/linear_att_cache_manager/config_objs.py index 6184437b9b..f588ec7d5c 100644 --- a/lightllm/common/linear_att_cache_manager/config_objs.py +++ b/lightllm/common/linear_att_cache_manager/config_objs.py @@ -68,9 +68,9 @@ def get_conv_state_shape(self): # Base committed sliding-window state, without speculative MTP tail. return (self.get_conv_dim(), self.conv_kernel_size - 1) - def get_spec_conv_state_shape(self, max_draft_step: int): + def get_mtp_conv_state_shape(self, mtp_step: int): # Working state with room for S speculative tokens before acceptance. - return (self.get_conv_dim(), (self.conv_kernel_size - 1) + max_draft_step) + return (self.get_conv_dim(), (self.conv_kernel_size - 1) + mtp_step) def get_ssm_state_shape(self): return (self.num_linear_v_heads, self.head_linear_k_dim, self.head_linear_v_dim) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 41328277b1..4a805cf113 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -7,7 +7,7 @@ from typing import List, Optional, TYPE_CHECKING from lightllm.common.basemodel.triton_kernel.gen_sampling_params import token_id_counter from lightllm.common.basemodel.triton_kernel.gen_sampling_params import update_req_to_token_id_counter -from lightllm.utils.envs_utils import get_env_start_args, enable_dynamic_spec +from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.config_utils import get_vocab_size from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.common.linear_att_cache_manager.layer_cache import LayerCache @@ -122,7 +122,9 @@ def __init__(self, max_request_num): device="cuda", ) self.req_to_next_token_probs = ( - torch.zeros_like(self.req_to_next_token_ids, dtype=torch.float32) if enable_dynamic_spec() else None + torch.zeros_like(self.req_to_next_token_ids, dtype=torch.float32) + if get_env_start_args().mtp_dynamic_verify + else None ) self.req_to_exponential_decay_length_penalty = torch.zeros( @@ -238,16 +240,16 @@ def gen_cpu_out_token_counter_sampling_params(self, req_objs: List["InferReq"]): class ReqManagerForMamba(ReqManager): def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_config: LinearAttCacheConfig): super().__init__(max_request_num, max_sequence_length, mem_manager) - self.max_draft_step = get_env_start_args().mtp_step + self.mtp_step = get_env_start_args().mtp_step # 因为在mtp的推理中,需要标记每个请求对应的mtp index状态(conv state 和 ssm state),在mtp对应序列中 # 的真实位置,所以需要需要一个标记来记录,不然算子无法找到真实的处理起点。 self.req_to_mtp_state_index = ( - torch.zeros((max_request_num + 1,), dtype=torch.int32, device="cuda") if self.max_draft_step > 0 else None + torch.zeros((max_request_num + 1,), dtype=torch.int32, device="cuda") if self.mtp_step > 0 else None ) # 突然想到, 在linear att 开启mtp的模式中,现在的prefill linear att 算子默认是从0的位置读取信息进行操作 # 所以不能支持 prefill decode mixed 操作了,因为一个decode过的请求,重新用prefill 算子跑,会出现读错linear # 状态位置的问题。导致bug, 在这里加个断言,以后可以支持上 TODO - if self.max_draft_step > 0: + if self.mtp_step > 0: assert get_env_start_args().enable_prefill_decode_mixed is False self.big_page_token_num = ( @@ -258,12 +260,12 @@ def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_con self.req_to_conv_state = LayerCache( size=(max_request_num + 1), dtype=self.linear_config.conv_state_dtype, - shape=self.linear_config.get_spec_conv_state_shape(max_draft_step=self.max_draft_step), + shape=self.linear_config.get_mtp_conv_state_shape(mtp_step=self.mtp_step), layer_num=self.linear_config.linear_layer_num, device="cuda", ) self.req_to_ssm_state = LayerCache( - size=(max_request_num + 1) * (self.max_draft_step + 1), + size=(max_request_num + 1) * (self.mtp_step + 1), dtype=self.linear_config.ssm_state_dtype, shape=self.linear_config.get_ssm_state_shape(), layer_num=self.linear_config.linear_layer_num, @@ -273,11 +275,11 @@ def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_con def init_linear_att_state(self, req: "InferReq"): conv_index = req.req_idx - ssm_start = req.req_idx * (self.max_draft_step + 1) + ssm_start = req.req_idx * (self.mtp_step + 1) self.req_to_conv_state.buffer[:, conv_index, ...].fill_(0) # #17: zero the FULL (mtp_step + 1)-row SSM block, not just canonical row +0, so a future # first-step verify reading offset>0 after fresh init never hits a never-written row (NaN). - self.req_to_ssm_state.buffer[:, ssm_start : ssm_start + (self.max_draft_step + 1), ...].fill_(0) + self.req_to_ssm_state.buffer[:, ssm_start : ssm_start + (self.mtp_step + 1), ...].fill_(0) if self.req_to_mtp_state_index is not None: self.req_to_mtp_state_index[req.req_idx] = 0 return @@ -298,7 +300,7 @@ def copy_big_page_buffer_to_linear_att_state(self, big_page_buffer_idx: int, req conv_state, ssm_state = big_page_buffers.get_state_cache(buffer_idx=big_page_buffer_idx) conv_dest = req.req_idx - ssm_dest = req.req_idx * (self.max_draft_step + 1) + ssm_dest = req.req_idx * (self.mtp_step + 1) conv_cache_width = conv_state.shape[-1] self.req_to_conv_state.buffer[:, conv_dest, ..., :conv_cache_width] = conv_state self.req_to_ssm_state.buffer[:, ssm_dest, ...] = ssm_state @@ -313,7 +315,7 @@ def copy_small_page_buffer_to_linear_att_state( buffer_idx=req.shared_kv_node.small_page_buffer_idx ) conv_dest = req.req_idx - ssm_dest = req.req_idx * (self.max_draft_step + 1) + ssm_dest = req.req_idx * (self.mtp_step + 1) conv_cache_width = conv_state.shape[-1] # TODO 下面这个从 cpu cache 拷贝数据的 gpu的操作,是否是阻塞的操作。 # 同时,非连续对象的拷贝,可能存在效率问题。 diff --git a/lightllm/models/__init__.py b/lightllm/models/__init__.py index e547068b1b..c1063e5bb3 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -54,57 +54,4 @@ from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel from lightllm.models.qwen3_eagle.model import Qwen3EagleModel from lightllm.models.qwen3_moe_mtp.model import Qwen3MOEMTPModel -from .registry import get_model, get_model_class - - -_ATTENTION_DRAFT_MODELS = { - "deepseek_v3": Deepseek3MTPModel, - "glm4_moe_lite": Glm4MoeLiteMTPModel, - "qwen3_5": Qwen3_5MTPModel, - "qwen3_5_text": Qwen3_5MTPModel, - "qwen3_5_moe": Qwen3_5MoeMTPModel, - "qwen3_5_moe_text": Qwen3_5MoeMTPModel, -} - -_NO_ATTENTION_DRAFT_MODELS = { - "mistral": MistralMTPModel, - "qwen3_moe": Qwen3MOEMTPModel, -} - - -def get_draft_model_class(model_cfg, spec_mode, is_linear_att_mixed_model=None): - architectures = set(model_cfg.get("architectures", ())) - - if spec_mode == "eagle3" and "Qwen3Eagle3Model" in architectures: - return Qwen3EagleModel - - if spec_mode == "dflash": - if "Qwen3_5DFlashModel" in architectures: - if is_linear_att_mixed_model is False: - raise ValueError("Qwen3_5DFlashModel requires a linear-attention mixed target") - return Qwen3_5DFlashModel - if architectures.intersection(("Qwen3DFlashModel", "Qwen3DSparkModel")): - if is_linear_att_mixed_model is True: - raise ValueError("linear-attention mixed targets require a Qwen3_5DFlashModel checkpoint") - return Qwen3DFlashModel - - if spec_mode == "dspark" and "Qwen3DSparkModel" in architectures: - if is_linear_att_mixed_model: - return Qwen3_5DSparkModel - return Qwen3DSparkModel - - model_type = model_cfg.get("model_type", "") - if model_type in _ATTENTION_DRAFT_MODELS: - if spec_mode not in ("vanilla_with_att", "eagle_with_att", "eagle3", "dspark", "dflash"): - raise ValueError(f"{model_type} requires an attention draft mode, got {spec_mode}") - return _ATTENTION_DRAFT_MODELS[model_type] - - if model_type in _NO_ATTENTION_DRAFT_MODELS: - if spec_mode not in ("vanilla_no_att", "eagle_no_att", "qwen3next_vanilla", "qwen3next_eagle"): - raise ValueError(f"{model_type} requires a no-attention draft mode, got {spec_mode}") - return _NO_ATTENTION_DRAFT_MODELS[model_type] - - raise ValueError( - f"Unsupported speculative draft model: mode={spec_mode}, " - f"model_type={model_cfg.get('model_type')}, architectures={sorted(architectures)}" - ) +from .registry import get_draft_model_class, get_model, get_model_class diff --git a/lightllm/models/deepseek_mtp/model.py b/lightllm/models/deepseek_mtp/model.py index e2b2a56137..959528c1e9 100644 --- a/lightllm/models/deepseek_mtp/model.py +++ b/lightllm/models/deepseek_mtp/model.py @@ -1,10 +1,15 @@ from typing import List from lightllm.models.deepseek2.model import Deepseek2TpPartModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.deepseek_mtp.layer_infer.pre_layer_infer import Deepseek3MTPPreLayerInfer from lightllm.models.deepseek_mtp.layer_weights.pre_and_post_layer_weight import Deepseek3MTPPreAndPostLayerWeight from lightllm.common.basemodel import TpPartBaseModel +@DraftModelRegistry( + model_type="deepseek_v3", + spec_modes=("vanilla_with_att", "eagle_with_att"), +) class Deepseek3MTPModel(Deepseek2TpPartModel): # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). diff --git a/lightllm/models/glm4_moe_lite_mtp/model.py b/lightllm/models/glm4_moe_lite_mtp/model.py index 2e4ba5c86b..3218378026 100644 --- a/lightllm/models/glm4_moe_lite_mtp/model.py +++ b/lightllm/models/glm4_moe_lite_mtp/model.py @@ -1,6 +1,7 @@ from typing import List from lightllm.models.deepseek_mtp.layer_infer.pre_layer_infer import Deepseek3MTPPreLayerInfer from lightllm.models.glm4_moe_lite.model import Glm4MoeLiteTpPartModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.glm4_moe_lite_mtp.layer_weights.pre_and_post_layer_weight import ( Glm4MoeLiteMTPPreAndPostLayerWeight, ) @@ -8,6 +9,10 @@ from lightllm.common.basemodel.basemodel import load_hf_weights +@DraftModelRegistry( + model_type="glm4_moe_lite", + spec_modes=("vanilla_with_att", "eagle_with_att"), +) class Glm4MoeLiteMTPModel(Glm4MoeLiteTpPartModel): # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). diff --git a/lightllm/models/mistral_mtp/model.py b/lightllm/models/mistral_mtp/model.py index f17bc0a383..ae1e208382 100644 --- a/lightllm/models/mistral_mtp/model.py +++ b/lightllm/models/mistral_mtp/model.py @@ -1,5 +1,6 @@ from typing import List from lightllm.models.mistral.model import MistralTpPartModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.mistral_mtp.layer_weights.pre_and_post_layer_weight import MistralMTPPreAndPostLayerWeight from lightllm.models.mistral_mtp.layer_infer.pre_layer_infer import MistralMTPPreLayerInfer from lightllm.models.mistral_mtp.layer_infer.post_layer_infer import MistralMTPPostLayerInfer @@ -8,6 +9,10 @@ from lightllm.common.basemodel import TpPartBaseModel +@DraftModelRegistry( + model_type="mistral", + spec_modes=("vanilla_no_att", "eagle_no_att"), +) class MistralMTPModel(MistralTpPartModel): # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). diff --git a/lightllm/models/qwen3_5_dflash/model.py b/lightllm/models/qwen3_5_dflash/model.py index 01ca9832f9..a9f3ce71cd 100644 --- a/lightllm/models/qwen3_5_dflash/model.py +++ b/lightllm/models/qwen3_5_dflash/model.py @@ -1,66 +1,49 @@ -from lightllm.distributed.communication_op import dist_group_manager from lightllm.models.llama.model import LlamaTpPartModel from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.qwen3_5_dflash.layer_weights.pre_and_post_layer_weight import ( Qwen35DFlashPreAndPostLayerWeight, ) +@DraftModelRegistry(model_type=("qwen3_5", "qwen3_5_text"), spec_modes="dflash") class Qwen3_5DFlashModel(Qwen3DFlashModel): """Qwen3.5 DFlash draft model.""" pre_and_post_weight_class = Qwen35DFlashPreAndPostLayerWeight - share_target_embedding_and_lm_head = True def _init_config(self): super()._init_config() - dflash_config = self.config.get("dflash_config", {}) - for key in ("target_layer_ids", "mask_token_id"): - if key not in self.config and key in dflash_config: - self.config[key] = dflash_config[key] + self.config.update(self.config.get("dflash_config", {})) - rope_parameters = self.config.get("rope_parameters") - if isinstance(rope_parameters, dict): - if "rope_theta" in rope_parameters and "rope_theta" not in self.config: - self.config["rope_theta"] = rope_parameters["rope_theta"] - if "partial_rotary_factor" in rope_parameters and "partial_rotary_factor" not in self.config: - self.config["partial_rotary_factor"] = rope_parameters["partial_rotary_factor"] - if "rope_scaling" not in self.config: - self.config["rope_scaling"] = rope_parameters + rope_parameters = self.config["rope_parameters"] + if "rope_theta" in rope_parameters and "rope_theta" not in self.config: + self.config["rope_theta"] = rope_parameters["rope_theta"] + if "partial_rotary_factor" in rope_parameters and "partial_rotary_factor" not in self.config: + self.config["partial_rotary_factor"] = rope_parameters["partial_rotary_factor"] + if "rope_scaling" not in self.config: + self.config["rope_scaling"] = rope_parameters def _init_custom(self): - # The Qwen3.5 target uses mrope/partial rotary with a different head - # shape from the ordinary DFlash draft. Build draft-owned rotary caches - # from the draft config instead of reusing main_model._cos/_sin. + # Draft and target use different rotary shapes, so the draft owns its rotary cache. LlamaTpPartModel._init_custom(self) - self.dist_group = dist_group_manager.get_default_group() - self.block_size = int(self.config["block_size"]) - self.mask_token_id = int(self.config["mask_token_id"]) + self.block_size = self.config["block_size"] + self.mask_token_id = self.config["mask_token_id"] def _init_mem_manager(self): target_mem_manager = self.main_model.mem_manager - target_linear_config = target_mem_manager.linear_config - - draft_head_dim = self.config.get("head_dim") - if draft_head_dim is None: - draft_head_dim = self.config["hidden_size"] // self.config["num_attention_heads"] - draft_kv_heads = int(self.config["num_key_value_heads"]) - target_kv_heads = int(target_linear_config.full_att_all_num_kv_heads) - draft_layer_num = int(self.config["n_layer"]) - reserved_draft_layer_num = int(target_linear_config.draft_full_att_kv_layer_num) - assert int(draft_head_dim) == int(target_mem_manager.head_dim) and draft_kv_heads == target_kv_heads, ( - "Qwen3.5 block draft currently requires draft and target full-attention KV shapes to match: " - f"draft=({draft_kv_heads}, {draft_head_dim}), " - f"target=({target_kv_heads}, {target_mem_manager.head_dim})" + draft_kv_shape = (self.config["num_key_value_heads"], self.config["head_dim"]) + target_kv_shape = ( + target_mem_manager.linear_config.full_att_all_num_kv_heads, + target_mem_manager.head_dim, ) - assert reserved_draft_layer_num >= draft_layer_num, ( - "Qwen3Next target KV cache did not reserve enough draft layers: " - f"required={draft_layer_num}, reserved={reserved_draft_layer_num}" + assert draft_kv_shape == target_kv_shape, ( + "Qwen3.5 block draft requires matching draft and target KV shapes, " + f"got draft={draft_kv_shape}, target={target_kv_shape}." ) - return super()._init_mem_manager() + super()._init_mem_manager() def _init_weights(self, start_layer_index=None): super()._init_weights(start_layer_index=start_layer_index) - if self.share_target_embedding_and_lm_head: - self.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ - self.pre_post_weight.lm_head_weight_ = self.main_model.pre_post_weight.lm_head_weight_ + self.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ + self.pre_post_weight.lm_head_weight_ = self.main_model.pre_post_weight.lm_head_weight_ diff --git a/lightllm/models/qwen3_5_dspark/model.py b/lightllm/models/qwen3_5_dspark/model.py index bf291d1a76..68cd359340 100644 --- a/lightllm/models/qwen3_5_dspark/model.py +++ b/lightllm/models/qwen3_5_dspark/model.py @@ -1,24 +1,39 @@ -from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel -from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import ( - Qwen3DSparkPostLayerInfer, -) -from lightllm.models.qwen3_dspark.model import DSparkModelOutputMixin -from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import ( - Qwen3DSparkPreAndPostLayerWeight, -) +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel +from lightllm.models.registry import DraftModelRegistry -class Qwen3_5DSparkModel(DSparkModelOutputMixin, Qwen3_5DFlashModel): - """DSpark draft model paired with a Qwen3.5 hybrid-attention target. +@DraftModelRegistry(model_type=("qwen3_5", "qwen3_5_text"), spec_modes="dspark") +class Qwen3_5DSparkModel(Qwen3DSparkModel): + """Qwen3 DSpark draft model paired with a Qwen3.5 target.""" - DeepSpec exports these checkpoints as ``Qwen3DSparkModel`` because the - draft backbone itself is a stack of Qwen3 full-attention layers. The - target pairing still matters at serving time: Qwen3.5 uses draft-owned - rotary caches and a target-owned compatible KV cache, while proposal logits - use DSpark's Markov/confidence heads and the ordinary zero-based block - layout. - """ + def _init_config(self): + super()._init_config() + self.config.update(self.config.get("dflash_config", {})) - pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight - post_layer_infer_class = Qwen3DSparkPostLayerInfer - share_target_embedding_and_lm_head = False + rope_parameters = self.config["rope_parameters"] + if "rope_theta" in rope_parameters and "rope_theta" not in self.config: + self.config["rope_theta"] = rope_parameters["rope_theta"] + if "partial_rotary_factor" in rope_parameters and "partial_rotary_factor" not in self.config: + self.config["partial_rotary_factor"] = rope_parameters["partial_rotary_factor"] + if "rope_scaling" not in self.config: + self.config["rope_scaling"] = rope_parameters + + def _init_custom(self): + # Draft and target use different rotary shapes, so the draft owns its rotary cache. + LlamaTpPartModel._init_custom(self) + self.block_size = self.config["block_size"] + self.mask_token_id = self.config["mask_token_id"] + + def _init_mem_manager(self): + target_mem_manager = self.main_model.mem_manager + draft_kv_shape = (self.config["num_key_value_heads"], self.config["head_dim"]) + target_kv_shape = ( + target_mem_manager.linear_config.full_att_all_num_kv_heads, + target_mem_manager.head_dim, + ) + assert draft_kv_shape == target_kv_shape, ( + "Qwen3.5 block draft requires matching draft and target KV shapes, " + f"got draft={draft_kv_shape}, target={target_kv_shape}." + ) + super()._init_mem_manager() diff --git a/lightllm/models/qwen3_5_moe_mtp/model.py b/lightllm/models/qwen3_5_moe_mtp/model.py index 022864f6b3..9e45501597 100644 --- a/lightllm/models/qwen3_5_moe_mtp/model.py +++ b/lightllm/models/qwen3_5_moe_mtp/model.py @@ -1,8 +1,13 @@ from lightllm.models.qwen3_5_mtp.model import Qwen3_5MTPModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.qwen3_5_moe_mtp.layer_weights.transformer_layer_weight import ( Qwen3_5MoeMTPTransformerLayerWeight, ) +@DraftModelRegistry( + model_type=("qwen3_5_moe", "qwen3_5_moe_text"), + spec_modes=("vanilla_with_att", "eagle_with_att"), +) class Qwen3_5MoeMTPModel(Qwen3_5MTPModel): transformer_weight_class = Qwen3_5MoeMTPTransformerLayerWeight diff --git a/lightllm/models/qwen3_5_mtp/model.py b/lightllm/models/qwen3_5_mtp/model.py index b0f55af7f3..760e08dd97 100644 --- a/lightllm/models/qwen3_5_mtp/model.py +++ b/lightllm/models/qwen3_5_mtp/model.py @@ -2,12 +2,17 @@ from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.models.qwen3_5.model import Qwen3_5TpPartModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.qwen3_5.layer_infer.transformer_layer_infer import Qwen35TransformerLayerInfer from lightllm.models.qwen3_5_mtp.layer_weights.pre_and_post_layer_weight import Qwen3_5MTPPreAndPostLayerWeight from lightllm.models.qwen3_5_mtp.layer_weights.transformer_layer_weight import Qwen3_5MTPTransformerLayerWeight from lightllm.models.qwen3_5_mtp.layer_infer.pre_layer_infer import Qwen3_5MTPPreLayerInfer +@DraftModelRegistry( + model_type=("qwen3_5", "qwen3_5_text"), + spec_modes=("vanilla_with_att", "eagle_with_att"), +) class Qwen3_5MTPModel(Qwen3_5TpPartModel): pre_and_post_weight_class = Qwen3_5MTPPreAndPostLayerWeight pre_layer_infer_class = Qwen3_5MTPPreLayerInfer diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py index a34d26a725..ba62d247a6 100644 --- a/lightllm/models/qwen3_dflash/model.py +++ b/lightllm/models/qwen3_dflash/model.py @@ -1,13 +1,10 @@ from lightllm.common.basemodel.attention import ( - BaseAttBackend, Fa3AttBackend, Fp8Fa3AttBackend, - get_decode_att_backend_class, - get_prefill_att_backend_class, ) from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.distributed.communication_op import dist_group_manager from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer from lightllm.models.qwen3_dflash.layer_infer.pre_layer_infer import Qwen3DFlashPreLayerInfer @@ -16,6 +13,7 @@ from lightllm.models.qwen3_dflash.layer_weights.transformer_layer_weight import Qwen3DFlashTransformerLayerWeight +@DraftModelRegistry(model_type="qwen3", spec_modes="dflash") class Qwen3DFlashModel(LlamaTpPartModel): """Qwen3 DFlash draft model.""" @@ -43,7 +41,6 @@ def _verify_params(self): def _init_custom(self): self._cos_cached = self.main_model._cos_cached self._sin_cached = self.main_model._sin_cached - self.dist_group = dist_group_manager.get_default_group() self.block_size = int(self.config["block_size"]) self.mask_token_id = int(self.config["mask_token_id"]) @@ -55,22 +52,11 @@ def _init_mem_manager(self): self.mem_manager = self.main_model.mem_manager def _init_att_backend(self): - self.prefill_att_backend: BaseAttBackend = get_prefill_att_backend_class(index=0)(model=self) - try: - self.decode_att_backend: BaseAttBackend = get_decode_att_backend_class( - index=0, - priority_list=["fa3"], - )(model=self) - except KeyError as exc: - raise NotImplementedError( - "Qwen3DFlashModel requires FA3 decode attention: " - "block draft attention is non-causal and Triton/FlashInfer decode paths do not honor decode_causal." - ) from exc + super()._init_att_backend() + # FA3 is currently the only backend that honors decode_causal=False. + # TODO: Remove this restriction after Triton and FlashInfer support non-causal block decode. if not isinstance(self.decode_att_backend, (Fa3AttBackend, Fp8Fa3AttBackend)): - raise NotImplementedError( - "Qwen3DFlashModel requires FA3 decode attention: " - "block draft attention is non-causal and Triton/FlashInfer decode paths do not honor decode_causal." - ) + raise NotImplementedError("Qwen3 DFlash decode requires FA3") def _init_infer_layer(self, start_layer_index=None): assert start_layer_index is None diff --git a/lightllm/models/qwen3_dspark/infer_struct.py b/lightllm/models/qwen3_dspark/infer_struct.py new file mode 100644 index 0000000000..0426bd87f7 --- /dev/null +++ b/lightllm/models/qwen3_dspark/infer_struct.py @@ -0,0 +1,8 @@ +from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo + + +class Qwen3DSparkInferStateInfo(Qwen3DFlashInferStateInfo): + def __init__(self): + super().__init__() + self.confidence_logits = None + self.draft_token_ids = None diff --git a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py index 5fca8a44b3..5134fadbd2 100644 --- a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py +++ b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py @@ -1,10 +1,8 @@ -import numpy as np import torch -import torch.nn.functional as F -from lightllm.distributed.communication_op import all_gather, all_gather_into_tensor -from lightllm.models.llama.infer_struct import LlamaInferStateInfo +from lightllm.distributed.communication_op import all_gather_into_tensor from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer +from lightllm.models.qwen3_dspark.infer_struct import Qwen3DSparkInferStateInfo from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import ( Qwen3DSparkPreAndPostLayerWeight, ) @@ -20,110 +18,45 @@ class Qwen3DSparkPostLayerInfer(Qwen3DFlashPostLayerInfer): def __init__(self, network_config): super().__init__(network_config) - self.block_size_ = int(network_config["block_size"]) - self.markov_rank_ = int(network_config.get("markov_rank", 0)) - self.markov_head_type_ = str(network_config.get("markov_head_type", "")).lower() - self.enable_confidence_head_ = bool(network_config.get("enable_confidence_head", False)) - self.confidence_head_with_markov_ = bool(network_config.get("confidence_head_with_markov", False)) - self.confidence_logits = None - self.draft_token_ids = None - - def pop_confidence_logits(self): - logits = self.confidence_logits - self.confidence_logits = None - return logits - - def pop_draft_token_ids(self): - token_ids = self.draft_token_ids - self.draft_token_ids = None - return token_ids - - def has_markov_head(self) -> bool: - return self.markov_rank_ > 0 - - def has_confidence_head(self) -> bool: - return self.enable_confidence_head_ - - def _linear_parameter(self, input_tensor: torch.Tensor, parameter_weight) -> torch.Tensor: - assert parameter_weight is not None - weight = parameter_weight.weight - bias = parameter_weight.bias - return F.linear(input_tensor.to(dtype=weight.dtype), weight, bias) + self.block_size_ = network_config["block_size"] + self.markov_rank_ = network_config.get("markov_rank", 0) + self.markov_head_type_ = (network_config.get("markov_head_type") or "").lower() + self.enable_confidence_head_ = network_config.get("enable_confidence_head", False) + self.confidence_head_with_markov_ = network_config.get("confidence_head_with_markov", False) def _markov_prev_embeddings( self, token_ids: torch.Tensor, layer_weight: Qwen3DSparkPreAndPostLayerWeight, ) -> torch.Tensor: - assert layer_weight.markov_w1_weight_ is not None - return F.embedding(token_ids.long(), layer_weight.markov_w1_weight_.weight) - - def _markov_project_bias( - self, - latent_states: torch.Tensor, - layer_weight: Qwen3DSparkPreAndPostLayerWeight, - ) -> torch.Tensor: - assert layer_weight.markov_w2_weight_ is not None - weight = layer_weight.markov_w2_weight_.weight - return F.linear(latent_states.to(dtype=weight.dtype), weight) + return torch.nn.functional.embedding(token_ids, layer_weight.markov_w1_weight_.weight) - def _markov_step_bias( + def _markov_step_latent( self, - prev_token_ids: torch.Tensor, + prev_embeddings: torch.Tensor, hidden_states: torch.Tensor, state: torch.Tensor, layer_weight: Qwen3DSparkPreAndPostLayerWeight, ): - prev_embeddings = self._markov_prev_embeddings(prev_token_ids, layer_weight) if self.markov_head_type_ == "vanilla": - return state, self._markov_project_bias(prev_embeddings, layer_weight) + return state, prev_embeddings hidden_states = hidden_states.to(dtype=prev_embeddings.dtype) if self.markov_head_type_ == "gated": gate_input = torch.cat([hidden_states, prev_embeddings], dim=-1) - gate = torch.sigmoid(self._linear_parameter(gate_input, layer_weight.markov_gate_proj_weight_)) - return state, self._markov_project_bias(gate * prev_embeddings, layer_weight) + gate = torch.sigmoid(layer_weight.markov_gate_proj_weight_.mm(gate_input)) + return state, gate * prev_embeddings assert self.markov_head_type_ == "rnn" if state is None: state = torch.zeros_like(prev_embeddings) joint_input = torch.cat([state, prev_embeddings, hidden_states], dim=-1) - joint = self._linear_parameter(joint_input, layer_weight.markov_joint_proj_weight_) + joint = layer_weight.markov_joint_proj_weight_.mm(joint_input) gate_raw, candidate_raw, output_raw = joint.chunk(3, dim=-1) gate = torch.sigmoid(gate_raw) candidate = torch.tanh(candidate_raw) state = gate * state + (1.0 - gate) * candidate - return state, self._markov_project_bias(torch.tanh(output_raw), layer_weight) - - @torch.no_grad() - def apply_markov_logits( - self, - base_logits: torch.Tensor, - block_hidden: torch.Tensor, - anchor_token_ids: torch.Tensor, - layer_weight: Qwen3DSparkPreAndPostLayerWeight, - ): - if not self.has_markov_head(): - return base_logits, torch.argmax(base_logits, dim=-1) - - sampled_tokens = [] - corrected_logits = [] - prev_token_ids = anchor_token_ids.long() - state = None - for step_idx in range(base_logits.shape[1]): - state, markov_bias = self._markov_step_bias( - prev_token_ids=prev_token_ids, - hidden_states=block_hidden[:, step_idx, :], - state=state, - layer_weight=layer_weight, - ) - step_logits = base_logits[:, step_idx, :] + markov_bias - next_token_ids = torch.argmax(step_logits, dim=-1) - sampled_tokens.append(next_token_ids) - corrected_logits.append(step_logits.unsqueeze(1)) - prev_token_ids = next_token_ids - - return torch.cat(corrected_logits, dim=1), torch.stack(sampled_tokens, dim=1) + return state, torch.tanh(output_raw) @torch.no_grad() def predict_confidence_logits( @@ -133,7 +66,7 @@ def predict_confidence_logits( sampled_tokens: torch.Tensor, layer_weight: Qwen3DSparkPreAndPostLayerWeight, ): - if not self.has_confidence_head(): + if not self.enable_confidence_head_: return None features = block_hidden @@ -145,96 +78,35 @@ def predict_confidence_logits( prev_embeddings = self._markov_prev_embeddings(prev_token_ids, layer_weight).to(dtype=block_hidden.dtype) features = torch.cat([block_hidden, prev_embeddings], dim=-1) - logits = self._linear_parameter(features, layer_weight.confidence_head_weight_) - return logits.float().squeeze(-1) + logits = layer_weight.confidence_head_weight_.mm(features.flatten(0, -2)) + return logits.float().view(features.shape[:-1]) - def _token_forward_with_hidden( - self, - input_embdings: torch.Tensor, - infer_state: LlamaInferStateInfo, - layer_weight: Qwen3DSparkPreAndPostLayerWeight, - ): - last_input, token_num = self._slice_get_last_input(input_embdings, infer_state) - input_embdings_dtype = input_embdings.dtype - head_hidden = last_input - normed_input = self._norm(last_input, infer_state, layer_weight) - lm_head_input = normed_input.permute(1, 0).reshape(-1, token_num) - logic_batch = layer_weight.lm_head_weight_(input=lm_head_input, alloc_func=self.alloc_tensor) - normed_input = None - lm_head_input = None - vocab_size = layer_weight.lm_head_weight_.vocab_size - if self.tp_world_size_ == 1: - gather_data = logic_batch - else: - gather_data = self.alloc_tensor((vocab_size, token_num), dtype=input_embdings_dtype) - split_indexes = np.linspace(0, vocab_size, self.tp_world_size_ + 1, dtype=np.int64) - all_gather( - [gather_data[split_indexes[i] : split_indexes[i + 1], :] for i in range(self.tp_world_size_)], - logic_batch, - group=infer_state.dist_group, - async_op=False, - ) - logic_batch = None - logits = self.alloc_tensor( - (token_num, vocab_size), - dtype=torch.float32, - ) - logits[:, :] = gather_data.permute(1, 0) - gather_data = None - return logits, head_hidden - - def _token_forward_with_local_logits_and_hidden( - self, - input_embdings: torch.Tensor, - infer_state: LlamaInferStateInfo, - layer_weight: Qwen3DSparkPreAndPostLayerWeight, - ): - """Project the LM head but keep its vocabulary shard local. - - Vanilla Markov decoding only needs the global maximum token at each - block position. Keeping logits sharded avoids a full-vocabulary TP - all-gather and prevents every rank from repeating the same Markov - projection over the complete vocabulary. - """ - - last_input, token_num = self._slice_get_last_input(input_embdings, infer_state) - head_hidden = last_input - normed_input = self._norm(last_input, infer_state, layer_weight) - lm_head_input = normed_input.permute(1, 0).reshape(-1, token_num) - local_logits = layer_weight.lm_head_weight_(input=lm_head_input, alloc_func=self.alloc_tensor) - return local_logits, head_hidden - - def _sample_tp_sharded_vanilla_markov( + def _sample_markov( self, local_logits: torch.Tensor, - infer_state: LlamaInferStateInfo, + block_hidden: torch.Tensor, + infer_state: Qwen3DSparkInferStateInfo, anchor_token_ids: torch.Tensor, layer_weight: Qwen3DSparkPreAndPostLayerWeight, ) -> torch.Tensor: - """Run exact greedy Markov decoding on TP vocabulary shards.""" - - assert self.markov_head_type_ == "vanilla" - assert layer_weight.markov_w2_weight_ is not None - token_num = local_logits.shape[1] - assert token_num % self.block_size_ == 0 - num_reqs = token_num // self.block_size_ - - vocab_size = int(layer_weight.lm_head_weight_.vocab_size) - split_indexes = np.linspace(0, vocab_size, self.tp_world_size_ + 1, dtype=np.int64) - local_start = int(split_indexes[self.tp_rank_]) - local_end = int(split_indexes[self.tp_rank_ + 1]) - assert local_logits.shape[0] == local_end - local_start, ( - f"local LM head rows must match TP vocabulary shard [{local_start}, {local_end}), " - f"got {local_logits.shape[0]}" - ) - local_markov_w2 = layer_weight.markov_w2_weight_.weight[local_start:local_end, :] + """Run sequential Markov decoding over TP-local vocabulary logits.""" - prev_token_ids = anchor_token_ids.long() + num_reqs = anchor_token_ids.shape[0] + local_start = layer_weight.markov_w2_weight_.tp_vocab_start_id + prev_token_ids = anchor_token_ids + state = None sampled_tokens = [] req_rows = torch.arange(num_reqs, dtype=torch.long, device=local_logits.device) for step_idx in range(self.block_size_): prev_embeddings = self._markov_prev_embeddings(prev_token_ids, layer_weight) - local_markov_bias = F.linear(prev_embeddings.to(dtype=local_markov_w2.dtype), local_markov_w2) + state, markov_latent = self._markov_step_latent( + prev_embeddings=prev_embeddings, + hidden_states=block_hidden[:, step_idx, :], + state=state, + layer_weight=layer_weight, + ) + markov_w2 = layer_weight.markov_w2_weight_.weight + local_markov_bias = torch.mm(markov_latent.to(dtype=markov_w2.dtype), markov_w2.t()) local_base_logits = local_logits[:, step_idx :: self.block_size_].permute(1, 0).float() local_scores = local_base_logits + local_markov_bias local_max_values, local_max_indexes = torch.max(local_scores, dim=-1) @@ -269,11 +141,9 @@ def _sample_tp_sharded_vanilla_markov( def token_forward( self, input_embdings: torch.Tensor, - infer_state: LlamaInferStateInfo, + infer_state: Qwen3DSparkInferStateInfo, layer_weight: Qwen3DSparkPreAndPostLayerWeight, ): - self.confidence_logits = None - self.draft_token_ids = None if infer_state.is_prefill: return super().token_forward( input_embdings=input_embdings, @@ -281,64 +151,39 @@ def token_forward( layer_weight=layer_weight, ) - use_tp_sharded_markov = ( - self.tp_world_size_ > 1 and self.has_markov_head() and self.markov_head_type_ == "vanilla" - ) - if use_tp_sharded_markov: - local_logits, head_hidden = self._token_forward_with_local_logits_and_hidden( - input_embdings=input_embdings, - infer_state=infer_state, - layer_weight=layer_weight, - ) - token_num = local_logits.shape[1] - assert token_num % self.block_size_ == 0 - num_reqs = token_num // self.block_size_ - block_hidden = head_hidden.reshape(num_reqs, self.block_size_, -1) - anchor_token_ids = infer_state.input_ids.reshape(num_reqs, self.block_size_)[:, 0] - sampled_tokens = self._sample_tp_sharded_vanilla_markov( + last_input, token_num = self._slice_get_last_input(input_embdings, infer_state) + num_reqs = token_num // self.block_size_ + block_hidden = last_input.reshape(num_reqs, self.block_size_, -1) + anchor_token_ids = infer_state.input_ids.reshape(num_reqs, self.block_size_)[:, 0] + + if self.markov_rank_ > 0: + normed_input = self._norm(last_input, infer_state, layer_weight) + lm_head_input = normed_input.permute(1, 0).reshape(-1, token_num) + local_logits = layer_weight.lm_head_weight_(input=lm_head_input, alloc_func=self.alloc_tensor) + sampled_tokens = self._sample_markov( local_logits, + block_hidden=block_hidden, infer_state=infer_state, anchor_token_ids=anchor_token_ids, layer_weight=layer_weight, ) - self.draft_token_ids = sampled_tokens.reshape(-1) - self.confidence_logits = self.predict_confidence_logits( + infer_state.draft_token_ids = sampled_tokens.reshape(-1) + infer_state.confidence_logits = self.predict_confidence_logits( block_hidden, anchor_token_ids=anchor_token_ids, sampled_tokens=sampled_tokens, layer_weight=layer_weight, ) - # The proposer consumes draft_token_ids directly. Keep the - # leading row dimension for generic graph padding/unpadding while - # avoiding an otherwise unused [rows, vocab] tensor. A single - # placeholder column is required because CUDA graph's no-ref - # tensor wrapper cannot represent a zero-byte allocation. + # Graph unpadding still uses the leading logits dimension when token ids are returned directly. return local_logits.new_empty((token_num, 1)) - logits, head_hidden = self._token_forward_with_hidden( - input_embdings=input_embdings, - infer_state=infer_state, - layer_weight=layer_weight, - ) - - assert ( - logits.shape[0] % self.block_size_ == 0 - ), f"DSpark draft logits rows must be a multiple of block_size={self.block_size_}, got {logits.shape[0]}" - num_reqs = logits.shape[0] // self.block_size_ + logits = self._lm_head_and_gather(last_input, token_num, layer_weight, infer_state) block_logits = logits.reshape(num_reqs, self.block_size_, -1) - block_hidden = head_hidden.reshape(num_reqs, self.block_size_, -1) - anchor_token_ids = infer_state.input_ids.reshape(num_reqs, self.block_size_)[:, 0] - - corrected_logits, sampled_tokens = self.apply_markov_logits( - block_logits, - block_hidden=block_hidden, - anchor_token_ids=anchor_token_ids, - layer_weight=layer_weight, - ) - self.confidence_logits = self.predict_confidence_logits( + sampled_tokens = torch.argmax(block_logits, dim=-1) + infer_state.confidence_logits = self.predict_confidence_logits( block_hidden, anchor_token_ids=anchor_token_ids, sampled_tokens=sampled_tokens, layer_weight=layer_weight, ) - return corrected_logits.reshape(logits.shape[0], -1) + return logits diff --git a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py index bf63a524ec..3965a8163a 100644 --- a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py +++ b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py @@ -1,4 +1,4 @@ -from lightllm.common.basemodel.layer_weights.meta_weights import ParameterWeight +from lightllm.common.basemodel.layer_weights.meta_weights import EmbeddingWeight, LMHeadWeight, ROWMMWeight from lightllm.common.quantization import Quantcfg from lightllm.models.qwen3_dflash.layer_weights.pre_and_post_layer_weight import Qwen3DFlashPreAndPostLayerWeight @@ -25,31 +25,42 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): self.markov_rank = markov_rank self.markov_head_type = str(network_config.get("markov_head_type", "")).lower() if markov_rank > 0: - self.markov_w1_weight_ = ParameterWeight( + # W1 is read once per Markov step; replication avoids a sequential TP collective on every lookup. + self.markov_w1_weight_ = EmbeddingWeight( + dim=markov_rank, + vocab_size=vocab_size, weight_name="markov_head.markov_w1.weight", data_type=self.data_type_, - weight_shape=(vocab_size, markov_rank), + tp_rank=0, + tp_world_size=1, ) - self.markov_w2_weight_ = ParameterWeight( + self.markov_w2_weight_ = LMHeadWeight( + dim=markov_rank, + vocab_size=vocab_size, weight_name="markov_head.markov_w2.weight", data_type=self.data_type_, - weight_shape=(vocab_size, markov_rank), ) if self.markov_head_type == "gated": - self.markov_gate_proj_weight_ = ParameterWeight( - weight_name="markov_head.gate_proj.weight", - bias_name="markov_head.gate_proj.bias", + self.markov_gate_proj_weight_ = ROWMMWeight( + in_dim=hidden_size + markov_rank, + out_dims=[markov_rank], + weight_names="markov_head.gate_proj.weight", + bias_names="markov_head.gate_proj.bias", data_type=self.data_type_, - weight_shape=(markov_rank, hidden_size + markov_rank), - bias_shape=(markov_rank,), + quant_method=self.quant_cfg.get_quant_method(0, "markov_head.gate_proj"), + tp_rank=0, + tp_world_size=1, ) elif self.markov_head_type == "rnn": - self.markov_joint_proj_weight_ = ParameterWeight( - weight_name="markov_head.joint_proj.weight", - bias_name="markov_head.joint_proj.bias", + self.markov_joint_proj_weight_ = ROWMMWeight( + in_dim=hidden_size + 2 * markov_rank, + out_dims=[3 * markov_rank], + weight_names="markov_head.joint_proj.weight", + bias_names="markov_head.joint_proj.bias", data_type=self.data_type_, - weight_shape=(3 * markov_rank, hidden_size + 2 * markov_rank), - bias_shape=(3 * markov_rank,), + quant_method=self.quant_cfg.get_quant_method(0, "markov_head.joint_proj"), + tp_rank=0, + tp_world_size=1, ) else: assert self.markov_head_type == "vanilla", f"unsupported DSpark markov head {self.markov_head_type}" @@ -57,10 +68,13 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): self.confidence_head_weight_ = None if enable_confidence_head: confidence_input_dim = hidden_size + (markov_rank if confidence_head_with_markov else 0) - self.confidence_head_weight_ = ParameterWeight( - weight_name="confidence_head.proj.weight", - bias_name="confidence_head.proj.bias", + self.confidence_head_weight_ = ROWMMWeight( + in_dim=confidence_input_dim, + out_dims=[1], + weight_names="confidence_head.proj.weight", + bias_names="confidence_head.proj.bias", data_type=self.data_type_, - weight_shape=(1, confidence_input_dim), - bias_shape=(1,), + quant_method=self.quant_cfg.get_quant_method(0, "confidence_head.proj"), + tp_rank=0, + tp_world_size=1, ) diff --git a/lightllm/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py index 48a870a66d..0b33fecfa3 100644 --- a/lightllm/models/qwen3_dspark/model.py +++ b/lightllm/models/qwen3_dspark/model.py @@ -1,26 +1,42 @@ from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel +from lightllm.models.registry import DraftModelRegistry +from lightllm.models.qwen3_dspark.infer_struct import Qwen3DSparkInferStateInfo from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import Qwen3DSparkPreAndPostLayerWeight from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor -class DSparkModelOutputMixin: - def _token_forward(self, infer_state): +@DraftModelRegistry(model_type="qwen3", spec_modes="dspark") +class Qwen3DSparkModel(Qwen3DFlashModel): + """Qwen3 DSpark draft model. + + DSpark keeps the DFlash block backbone. Its extra Markov/confidence heads + run in the post layer so the proposer can consume corrected logits through + the same token-generation path as DFlash. + """ + + pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight + post_layer_infer_class = Qwen3DSparkPostLayerInfer + infer_state_class = Qwen3DSparkInferStateInfo + + def _verify_params(self): + super()._verify_params() + assert self.config.get("enable_confidence_head", False), "DSpark requires enable_confidence_head=true" + + def _token_forward(self, infer_state: Qwen3DSparkInferStateInfo): model_output = super()._token_forward(infer_state) - confidence_logits = self.post_infer.pop_confidence_logits() - draft_token_ids = self.post_infer.pop_draft_token_ids() if infer_state.is_cuda_graph: - if confidence_logits is not None: - confidence_logits = tensor_to_no_ref_tensor(confidence_logits) - if draft_token_ids is not None: - draft_token_ids = tensor_to_no_ref_tensor(draft_token_ids) + if infer_state.confidence_logits is not None: + infer_state.confidence_logits = tensor_to_no_ref_tensor(infer_state.confidence_logits) + if infer_state.draft_token_ids is not None: + infer_state.draft_token_ids = tensor_to_no_ref_tensor(infer_state.draft_token_ids) return DSparkModelOutput( logits=model_output.logits, spec_hidden=model_output.spec_hidden, - confidence_logits=confidence_logits, - draft_token_ids=draft_token_ids, + confidence_logits=infer_state.confidence_logits, + draft_token_ids=infer_state.draft_token_ids, ) def _create_unpad_decode_model_output(self, model_output: DSparkModelOutput, origin_batch_size: int): @@ -38,19 +54,3 @@ def _create_unpad_decode_model_output(self, model_output: DSparkModelOutput, ori assert origin_batch_size % rows_per_confidence == 0 model_output.confidence_logits = model_output.confidence_logits[: origin_batch_size // rows_per_confidence] return model_output - - -class Qwen3DSparkModel(DSparkModelOutputMixin, Qwen3DFlashModel): - """Qwen3 DSpark draft model. - - DSpark keeps the DFlash block backbone. Its extra Markov/confidence heads - run in the post layer so the proposer can consume corrected logits through - the same token-generation path as DFlash. - """ - - pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight - post_layer_infer_class = Qwen3DSparkPostLayerInfer - - def _verify_params(self): - super()._verify_params() - assert self.config.get("enable_confidence_head", False), "DSpark requires enable_confidence_head=true" diff --git a/lightllm/models/qwen3_eagle/model.py b/lightllm/models/qwen3_eagle/model.py index 93c6deddfd..8f54373023 100644 --- a/lightllm/models/qwen3_eagle/model.py +++ b/lightllm/models/qwen3_eagle/model.py @@ -3,12 +3,14 @@ from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.qwen3_eagle.layer_infer.pre_layer_infer import Qwen3EaglePreLayerInfer from lightllm.models.qwen3_eagle.layer_infer.transformer_layer_infer import Qwen3EagleTransformerLayerInfer from lightllm.models.qwen3_eagle.layer_weights.pre_and_post_layer_weight import Qwen3EaglePreAndPostLayerWeight from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight +@DraftModelRegistry(model_type="qwen3", spec_modes="eagle3") class Qwen3EagleModel(LlamaTpPartModel): is_mtp_draft_model = True diff --git a/lightllm/models/qwen3_moe_mtp/model.py b/lightllm/models/qwen3_moe_mtp/model.py index d9854250e2..a4b95a21fc 100644 --- a/lightllm/models/qwen3_moe_mtp/model.py +++ b/lightllm/models/qwen3_moe_mtp/model.py @@ -1,5 +1,6 @@ from typing import List from lightllm.models.qwen3_moe.model import Qwen3MOEModel +from lightllm.models.registry import DraftModelRegistry from lightllm.models.qwen3_moe_mtp.layer_weights.pre_and_post_layer_weight import Qwen3MOEMTPPreAndPostLayerWeight from lightllm.models.deepseek_mtp.layer_infer.pre_layer_infer import Deepseek3MTPPreLayerInfer from lightllm.models.qwen3_moe_mtp.layer_infer.transformer_layer_infer import Qwen3MOEMTPTransformerLayerInfer @@ -7,6 +8,10 @@ from lightllm.common.basemodel import TpPartBaseModel +@DraftModelRegistry( + model_type="qwen3_moe", + spec_modes=("vanilla_no_att", "eagle_no_att"), +) class Qwen3MOEMTPModel(Qwen3MOEModel): # MTP draft model marker (consumed by the decode CUDA-graph / padding paths). diff --git a/lightllm/models/registry.py b/lightllm/models/registry.py index e9568cc3d8..b2e7ffa4ac 100644 --- a/lightllm/models/registry.py +++ b/lightllm/models/registry.py @@ -1,10 +1,7 @@ import collections -from lightllm.utils.log_utils import init_logger - -logger = init_logger(__name__) - from dataclasses import dataclass -from typing import Type, Dict, Optional, Callable, List, Union, TypeVar +from typing import Callable, Dict, List, Optional, Tuple, Type, TypeVar, Union + from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) @@ -92,6 +89,42 @@ def get_model_class(self, model_cfg: dict): ModelRegistry = _ModelRegistries() +class _DraftModelRegistry: + def __init__(self): + self._registry: Dict[Tuple[str, str], Type] = {} + + def __call__( + self, + model_type: Union[str, List[str], Tuple[str, ...]], + spec_modes: Union[str, List[str], Tuple[str, ...]], + ) -> Callable[[T], T]: + model_types = (model_type,) if isinstance(model_type, str) else tuple(model_type) + modes = (spec_modes,) if isinstance(spec_modes, str) else tuple(spec_modes) + + def decorator(model_class: T) -> T: + for current_model_type in model_types: + for spec_mode in modes: + key = (current_model_type, spec_mode) + if key in self._registry: + raise ValueError(f"Duplicate draft model registration: {key}") + self._registry[key] = model_class + return model_class + + return decorator + + def get_model_class(self, model_cfg: dict, spec_mode: str) -> Type: + model_type = model_cfg.get("model_type", "") + try: + return self._registry[(model_type, spec_mode)] + except KeyError: + raise ValueError( + f"Unsupported speculative draft model: mode={spec_mode}, model_type={model_type}" + ) from None + + +DraftModelRegistry = _DraftModelRegistry() + + def get_model(model_cfg: dict, model_kvargs: dict): try: model, is_multimodal = ModelRegistry.get_model(model_cfg, model_kvargs) @@ -110,6 +143,10 @@ def get_model_class(model_cfg: dict): raise +def get_draft_model_class(model_cfg, spec_mode): + return DraftModelRegistry.get_model_class(model_cfg=model_cfg, spec_mode=spec_mode) + + def is_reward_model() -> Callable[[Dict[str, any]], bool]: """Predicate: whether the model is RewardModel.""" return lambda model_cfg: "RewardModel" in model_cfg.get("architectures", [""])[0] diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 515cacf699..b4ebd56906 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -734,8 +734,6 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "vanilla_no_att", "eagle_no_att", "eagle3", - "qwen3next_vanilla", - "qwen3next_eagle", "dspark", "dflash", None, diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 16fcea1525..a689b42713 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -162,31 +162,18 @@ def _launch_subprocesses(args: StartArgs): ) # mtp params check - spec_mode = args.mtp_mode - if spec_mode is not None: - if spec_mode in ("vanilla_with_att", "vanilla_no_att", "qwen3next_vanilla"): - assert args.mtp_step > 0 - draft_model_count = args.mtp_step - elif spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3", "qwen3next_eagle"): - assert args.mtp_step > 0 - draft_model_count = 1 - else: - assert spec_mode in ("dspark", "dflash"), f"unsupported speculative mode {spec_mode}" - assert args.mtp_step > 0 - draft_model_count = 1 - - if spec_mode == "dspark": - args.mtp_dynamic_verify = True - elif spec_mode == "dflash": - args.mtp_dynamic_verify = False - + if args.mtp_mode is not None: if args.mtp_draft_model_dir is None: - assert spec_mode not in ("dspark", "dflash"), f"--mtp_draft_model_dir is required for {spec_mode} mode" - args.mtp_draft_model_dir = [args.model_dir] * draft_model_count - assert len(args.mtp_draft_model_dir) >= draft_model_count + assert args.mtp_mode not in ( + "eagle3", + "dspark", + "dflash", + ), f"--mtp_draft_model_dir is required for {args.mtp_mode} mode" + args.mtp_draft_model_dir = [args.model_dir] * args.mtp_step + assert args.mtp_step > 0 else: - assert args.mtp_step == 0 assert args.mtp_draft_model_dir is None + assert args.mtp_step == 0 # automatically set visual_dp based on visual_tp and tp. # In visual proxy mode keep the caller-provided visual_dp / visual_tp. diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index af723753a0..245cb99029 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -190,8 +190,6 @@ class StartArgs: "eagle3", "dspark", "dflash", - "qwen3next_vanilla", - "qwen3next_eagle", None, ] }, diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 56d1c1445e..976a20f65b 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -566,11 +566,11 @@ def __init__( # 卸载到 cpu cache 中,该标志变量用于标记请求的卸载任务的状态 self.cpu_cache_task_status: "InferReq._CpuCacheTaskStatus" = InferReq._CpuCacheTaskStatus.NOT_STARTED - # max_draft_step 用来记录一个请求 draft模型每步需要生成的token数量 + # mtp_step 用来记录一个请求 draft模型每步需要生成的token数量 # 正常模式下,这个值为0,在 mtp 模式下,这个值为 draft 模型每步需要生成的token数量 - self.max_draft_step: int = get_env_start_args().mtp_step - if self.max_draft_step > 0: - self.decode_need_token_num = self._spec_decode_need_token_num + self.mtp_step: int = get_env_start_args().mtp_step + if self.mtp_step > 0: + self.decode_need_token_num = self._mtp_decode_need_token_num else: self.decode_need_token_num = self._normal_decode_need_token_num @@ -923,8 +923,8 @@ def decode_need_token_num(self) -> int: def _normal_decode_need_token_num(self) -> int: return 1 - def _spec_decode_need_token_num(self) -> int: - return (1 + self.max_draft_step) * 2 + def _mtp_decode_need_token_num(self) -> int: + return (1 + self.mtp_step) * 2 class InferReqUpdatePack: 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 837cf177fe..64ba873d81 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -107,13 +107,6 @@ def init_model(self, kvargs): self.is_multinode_tp = self.args.nnodes > 1 and self.args.dp == 1 self.is_pd_mode = self.run_mode in ["prefill", "decode"] self.is_pd_decode_mode = self.run_mode == "decode" - if self.args.mtp_mode in ("eagle3", "dspark", "dflash"): - assert ( - not self.args.enable_decode_microbatch_overlap - ), f"{self.args.mtp_mode} mode does not support decode microbatch overlap" - assert ( - not self.args.enable_prefill_microbatch_overlap - ), f"{self.args.mtp_mode} mode does not support prefill microbatch overlap" self.logger = init_logger(__name__) @@ -154,7 +147,6 @@ def init_model(self, kvargs): "quant_cfg": kvargs.get("quant_cfg", None), "expert_dtype": kvargs.get("expert_dtype", None), "run_mode": self.run_mode, - "hidden_layer_ids": self._target_hidden_layer_ids(model_cfg), "decode_batch_multiplier": target_decode_batch_multiplier, } self.model, self.is_multimodal = get_model(model_cfg, model_kvargs) @@ -254,10 +246,7 @@ def init_model(self, kvargs): # 只会在 pd pd 模式下才会使用,用于上传分块传输任务是否成功。 self.shm_pd_trans_io_buffer = ShmObjsIOBuffer(tail_str="pd") - # 开启 mtp 模式,需要完成mtp model的初始化 - if self.args.mtp_mode is not None: - self.init_spec_draft_model(model_kvargs) - self.spec_engine = SpecEngine(backend=self) + self.init_spec_engine(model_kvargs) if self.args.enable_cpu_cache: self.multi_level_cache_module = MultiLevelKvCacheModule(self) @@ -310,12 +299,18 @@ def prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): def decode(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): raise NotImplementedError() - def init_spec_draft_model(self, main_kvargs: dict): + def init_spec_engine(self, main_kvargs: dict): + if self.args.mtp_mode is None: + return + self.init_mtp_draft_model(main_kvargs) + self.spec_engine = SpecEngine(backend=self) + return + + def init_mtp_draft_model(self, main_kvargs: dict): self.max_draft_step = self.args.mtp_step self.draft_models = [] spec_mode = self.args.mtp_mode - is_chained_draft = spec_mode in ("vanilla_with_att", "vanilla_no_att", "qwen3next_vanilla") - is_recurrent_draft = spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3", "qwen3next_eagle") + is_chained_draft = spec_mode in ("vanilla_with_att", "vanilla_no_att") os.environ["DISABLE_CHECK_MAX_LEN_INFER"] = "1" @@ -328,11 +323,11 @@ def init_spec_draft_model(self, main_kvargs: dict): draft_model_cfg, _ = PretrainedConfig.get_config_dict(draft_model_dirs[i]) if is_chained_draft: draft_decode_batch_multiplier = self.max_draft_step + 1 - elif is_recurrent_draft: - draft_decode_batch_multiplier = 1 - else: + elif spec_mode in ("dspark", "dflash"): block_size = int(draft_model_cfg["block_size"]) draft_decode_batch_multiplier = block_size + else: + draft_decode_batch_multiplier = 1 draft_model_kvargs = { "weight_dir": draft_model_dirs[i], "max_total_token_num": self.model.mem_manager.size, @@ -360,7 +355,6 @@ def init_spec_draft_model(self, main_kvargs: dict): draft_model_class = get_draft_model_class( model_cfg=draft_model_cfg, spec_mode=spec_mode, - is_linear_att_mixed_model=self.is_linear_att_mixed_model, ) self.draft_models.append(draft_model_class(draft_model_kvargs)) @@ -469,28 +463,6 @@ def _capture_prompt_logprobs_if_needed( start_loc += q_len return - def _target_hidden_layer_ids(self, model_cfg: dict): - if self.args.mtp_mode not in ("eagle3", "dspark", "dflash"): - return None - - model_cfg = model_cfg.get("text_config", model_cfg) - layer_num = int(model_cfg.get("num_hidden_layers", model_cfg.get("n_layer"))) - - draft_model_dirs = self.args.mtp_draft_model_dir - assert draft_model_dirs - draft_cfg, _ = PretrainedConfig.get_config_dict(draft_model_dirs[0]) - target_layer_ids = draft_cfg.get("target_layer_ids") - if target_layer_ids is None and self.args.mtp_mode == "dflash": - target_layer_ids = draft_cfg.get("dflash_config", {}).get("target_layer_ids") - if target_layer_ids is None: - target_layer_ids = [1, layer_num // 2 - 1, layer_num - 4] - - target_layer_ids = tuple(int(layer_id) for layer_id in target_layer_ids) - assert target_layer_ids and all( - 0 <= layer_id < layer_num for layer_id in target_layer_ids - ), f"invalid target_layer_ids={target_layer_ids} for target layer_num={layer_num}" - return target_layer_ids - def _try_read_new_reqs(self): if self.is_multinode_tp: self._try_read_new_reqs_multinode_tp() @@ -849,13 +821,6 @@ def _post_handle( extra_post_req_handle_func 用于提供在一个请求确定输出的时候,给出额外的后处理操作,主要是用于 约束输出等模式,设置自己请求内部的状态机的状态,并添加额外的停止判定条件等。 """ - if isinstance(next_token_ids, torch.Tensor): - next_token_ids = next_token_ids.numpy() - if isinstance(next_token_logprobs, torch.Tensor): - next_token_logprobs = next_token_logprobs.numpy() - if isinstance(next_token_ranks, torch.Tensor): - next_token_ranks = next_token_ranks.numpy() - for req_obj, next_token_id, next_token_logprob, next_token_rank, pack in zip( run_reqs, next_token_ids, next_token_logprobs, next_token_ranks, run_reqs_update_packs ): @@ -905,7 +870,7 @@ def _update_spec_verify_token_num( if selected_run_reqs is None: for req in decode_reqs: - req.update_spec_verify_token_num(verify_token_num=req.max_draft_step + 1) + req.update_spec_verify_token_num(verify_token_num=req.mtp_step + 1) req.update_spec_verify_step_num(verify_step_num=1) return @@ -918,22 +883,12 @@ def _update_spec_verify_token_num( def _gen_argmax_token_ids(self, model_output: ModelOutput): logits = model_output.logits - draft_next_token_ids_gpu = torch.argmax(logits, dim=-1) - - # 如果draft和target的词表不同,需要把draft token映射回主模型词表。 - if self.args.mtp_mode == "eagle3": - draft_next_token_ids_gpu = self.draft_models[0].map_draft_vocab_to_main_vocab(draft_next_token_ids_gpu) - return draft_next_token_ids_gpu + return torch.argmax(logits, dim=-1) def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): logits = model_output.logits probs = torch.softmax(logits, dim=-1) max_probs, draft_next_token_ids_gpu = torch.max(probs, dim=-1) - - # 如果self.d2t不为None,那么draft的token需要进行相应的转换 - if self.args.mtp_mode == "eagle3": - draft_next_token_ids_gpu = self.draft_models[0].map_draft_vocab_to_main_vocab(draft_next_token_ids_gpu) - return draft_next_token_ids_gpu, max_probs def _sample_and_scatter_token( diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index cbffdbde9a..3f8c3c48d5 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -15,7 +15,7 @@ from lightllm.utils.dist_utils import get_current_device_id from .control_state import ControlState from lightllm.utils.dist_utils import create_new_group_for_current_dp -from lightllm.utils.envs_utils import enable_dynamic_spec, get_env_start_args +from lightllm.utils.envs_utils import get_env_start_args logger = init_logger(__name__) @@ -30,9 +30,9 @@ def __init__(self) -> None: # 在 mtp 模式下切换绑定的prefill 和 decode 函数 if get_env_start_args().mtp_mode is not None: - self.prefill = self.prefill_spec - self.decode = self.decode_spec - self.enable_dynamic_spec = enable_dynamic_spec() + self.prefill = self.prefill_mtp + self.decode = self.decode_mtp + self.enable_dynamic_spec = get_env_start_args().mtp_dynamic_verify else: self.prefill = self.prefill_normal self.decode = self.decode_normal @@ -40,10 +40,6 @@ def __init__(self) -> None: self.classed_req_strict_prefill = False return - # cpu 把算子提交到gpu 上 - # GPU - # CPU - def init_custom(self): super().init_custom() if self.enable_dynamic_spec: @@ -185,7 +181,7 @@ def decode_normal( event_pack.notify_pre_post_handle() return - def prefill_spec( + def prefill_mtp( self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq], @@ -244,7 +240,7 @@ def prefill_spec( event_pack.notify_pre_post_handle() return - def decode_spec( + def decode_mtp( self, event_pack: OverlapEventPack, decode_reqs: List[InferReq], diff --git a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py index 8c4dfc1a33..4802b83235 100644 --- a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py @@ -21,17 +21,15 @@ class DiversehBackend(ChunkedPrefillBackend): def __init__(self) -> None: super().__init__() - if get_env_start_args().mtp_mode: - # 当前只有 mistral 和 Qwen3Next mtp 可以使用 diverse mode 的 mtp 功能。 - self.prefill = self.beam_prefill - assert get_env_start_args().mtp_mode in [ + self.prefill = self.beam_prefill + spec_mode = get_env_start_args().mtp_mode + if spec_mode is not None: + assert spec_mode in [ + "vanilla_with_att", + "eagle_with_att", "vanilla_no_att", "eagle_no_att", - "qwen3next_vanilla", - "qwen3next_eagle", ] - else: - self.prefill = self.beam_prefill self.classed_req_strict_prefill = True diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index be87b8725b..458bf07e0b 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -37,21 +37,20 @@ def __init__(self) -> None: "eagle_with_att", "eagle_no_att", "eagle3", - "qwen3next_eagle", ) if self.enable_prefill_microbatch_overlap: - self.prefill = self.prefill_overlap_spec + self.prefill = self.prefill_overlap_mtp else: - self.prefill = self.prefill_spec + self.prefill = self.prefill_mtp if self.enable_decode_microbatch_overlap: - self.decode = self.decode_overlap_spec + self.decode = self.decode_overlap_mtp self._draft_decode_overlap_func = ( self._draft_decode_eagle_overlap if self.uses_recurrent_draft else self._draft_decode_vanilla_overlap ) else: - self.decode = self.decode_spec + self.decode = self.decode_mtp self._draft_decode_func = ( self._draft_decode_eagle if self.uses_recurrent_draft else self._draft_decode_vanilla ) @@ -430,7 +429,7 @@ def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferRe event_pack.notify_pre_post_handle() return - def prefill_spec(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): + def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): # main model prefill model_input, run_reqs, _ = padded_prepare_prefill_inputs(prefill_reqs) req_num = len(run_reqs) @@ -502,7 +501,7 @@ def prefill_spec(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq event_pack.notify_pre_post_handle() return - def decode_spec(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): + def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): model_input, run_reqs, _ = padded_prepare_decode_inputs(decode_reqs) b_mtp_index_cpu = model_input.b_mtp_index req_num = len(run_reqs) @@ -717,7 +716,7 @@ def _draft_decode_eagle( ) return proposal.extra_mem_indexes_cpu - def prefill_overlap_spec(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): + def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): ( model_input0, run_reqs0, @@ -811,7 +810,7 @@ def prefill_overlap_spec(self, event_pack: OverlapEventPack, prefill_reqs: List[ event_pack.notify_pre_post_handle() return - def decode_overlap_spec(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): + def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): ( model_input0, run_reqs0, diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py index d35737e95a..8221185572 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py @@ -8,7 +8,6 @@ from lightllm.utils.infer_utils import calculate_time from lightllm.utils.envs_utils import ( enable_diverse_mode_gqa_decode_fast_kernel, - enable_dynamic_spec, enable_triton_mtp_kernel, get_env_start_args, ) @@ -159,7 +158,7 @@ def padded_prepare_decode_inputs( b_mtp_index = [] b_seq_len = [] b_q_seq_len = [] - draft_step = get_env_start_args().mtp_step + args_mtp_step = get_env_start_args().mtp_step batch_multimodal_params = [] for req in req_objs: run_reqs.append(req) @@ -172,7 +171,7 @@ def padded_prepare_decode_inputs( b_mtp_index.append(0) batch_multimodal_params.append(req.multimodal_params) # process the draft tokens. - for step in range(req.max_draft_step): + for step in range(req.mtp_step): run_reqs.append(req) seq_len += 1 total_token_num += seq_len @@ -191,7 +190,7 @@ def padded_prepare_decode_inputs( b_q_seq_len.append(1) b_mtp_index.append(0) batch_multimodal_params.append({"images": [], "audios": []}) - for step in range(draft_step): + for step in range(args_mtp_step): seq_len += 1 total_token_num += seq_len b_seq_len.append(seq_len) @@ -208,13 +207,13 @@ def padded_prepare_decode_inputs( b_mtp_index = torch.tensor(b_mtp_index, dtype=torch.int32, device="cpu") b_position_delta = build_b_position_delta(batch_multimodal_params) - padded_row_count = padded_req_num * (draft_step + 1) + padded_row_count = padded_req_num * (args_mtp_step + 1) if enable_diverse_mode_gqa_decode_fast_kernel(): b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) if padded_row_count > 0: b_shared_seq_len = F.pad(b_shared_seq_len, (0, padded_row_count), value=0) b_mark_shared_group = F.pad(b_mark_shared_group, (0, padded_row_count), value=1) - elif enable_dynamic_spec() or enable_triton_mtp_kernel(): + elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): b_shared_seq_len = None b_mark_shared_group = build_spec_shared_group_markers(b_mtp_index=b_mtp_index) else: @@ -247,7 +246,7 @@ def padded_prepare_decode_inputs( b_position_delta=b_position_delta, b_shared_seq_len=b_shared_seq_len, b_mark_shared_group=b_mark_shared_group, - draft_step=draft_step, + draft_step=args_mtp_step, is_prefill=False, multimodal_params=batch_multimodal_params, ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py index d0e035aed0..1157275afb 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py @@ -30,9 +30,9 @@ def sample( skip_top_p, exist_req_use_random_seed, ) = _get_post_sample_tensors(reqs) + eos_ids = g_pin_mem_manager.gen_from_list(key="eos_ids", data=eos_id, dtype=torch.int32).cuda(non_blocking=True) sampling_params_manager = g_infer_context.req_manager.req_sampling_params_manager - sample_reqs = reqs if selected_row_mask is not None: ( b_req_idx, @@ -56,15 +56,11 @@ def sample( or exist_req_use_random_seed or sampling_params_manager.penalty_counter_mode == "cpu_counter" ): - sample_reqs = _get_selected_reqs(reqs=reqs, selected_row_mask=selected_row_mask) + reqs = _get_selected_reqs(reqs=reqs, selected_row_mask=selected_row_mask) if has_invalid_token_ids: - invalid_token_ids, cu_invalid_token_num, has_invalid_token_ids = _get_invalid_token_tensors( - reqs=sample_reqs - ) + invalid_token_ids, cu_invalid_token_num, has_invalid_token_ids = _get_invalid_token_tensors(reqs=reqs) if exist_req_use_random_seed: - exist_req_use_random_seed = any(req.generator is not None for req in sample_reqs) - - eos_ids = g_pin_mem_manager.gen_from_list(key="eos_ids", data=eos_id, dtype=torch.int32).cuda(non_blocking=True) + exist_req_use_random_seed = any(req.generator is not None for req in reqs) # 这里需要区分历史token的频率惩罚类的系数的生效模式,目前支持两种在线统计方式: # 一种是基于 cpu 的,每个 req 对象利用其上绑定的dict对象out_token_id_count,每生成一个token就进行相应 @@ -83,7 +79,7 @@ def sample( p_token_ids, p_token_counts, p_cumsum_seq_len, - ) = sampling_params_manager.gen_cpu_out_token_counter_sampling_params(req_objs=sample_reqs) + ) = sampling_params_manager.gen_cpu_out_token_counter_sampling_params(req_objs=reqs) apply_penalty( Logits=logits, @@ -123,13 +119,13 @@ def sample( elif skip_top_k and skip_top_p: # topk 等于整个词表,topp 等于1.0,等价于不进行topk topp过滤,直接进行随机采样,可以提升采样速度 - batch_next_token_ids = _random_sample(probs, sample_reqs, exist_req_use_random_seed) + batch_next_token_ids = _random_sample(probs, reqs, exist_req_use_random_seed) batch_next_token_probs = torch.gather(probs, dim=1, index=batch_next_token_ids.view(-1, 1)) return batch_next_token_ids.view(-1), torch.log(batch_next_token_probs).view(-1) else: batch_next_token_ids, batch_next_token_logprobs = _top_p_top_k_sample( - sample_reqs, probs, b_top_ps, b_top_ks, exist_req_use_random_seed + reqs, probs, b_top_ps, b_top_ks, exist_req_use_random_seed ) return batch_next_token_ids.view(-1), batch_next_token_logprobs.view(-1) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index e031b5e318..8daa4651d9 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -7,7 +7,6 @@ enable_diverse_mode_gqa_decode_fast_kernel, enable_triton_mtp_kernel, get_diverse_max_batch_shared_group_size, - enable_dynamic_spec, get_env_start_args, ) @@ -117,7 +116,7 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_mtp_index.append(0) multimodal_params.append(req.multimodal_params) # process the draft tokens. - for step in range(req.max_draft_step): + for step in range(req.mtp_step): run_reqs.append(req) b_req_idx.append(req.req_idx) seq_len += 1 @@ -137,7 +136,7 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In if enable_diverse_mode_gqa_decode_fast_kernel(): b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) - elif enable_dynamic_spec() or enable_triton_mtp_kernel(): + elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): b_shared_seq_len = None b_mark_shared_group = build_spec_shared_group_markers(b_mtp_index=b_mtp_index) else: diff --git a/lightllm/server/router/model_infer/speculative/engine.py b/lightllm/server/router/model_infer/speculative/engine.py index d4130641a5..f7dcd3c0f3 100644 --- a/lightllm/server/router/model_infer/speculative/engine.py +++ b/lightllm/server/router/model_infer/speculative/engine.py @@ -21,7 +21,6 @@ SpecDecodeRunner, ) from lightllm.server.router.model_infer.speculative.verifier import SpecVerifier, SpecVerifyResult -from lightllm.utils.envs_utils import enable_dynamic_spec class SpecEngine: @@ -42,7 +41,7 @@ class SpecEngine: def __init__(self, backend) -> None: self.backend = backend self.spec_mode = backend.args.mtp_mode - self.enable_dynamic_spec = enable_dynamic_spec() + self.enable_dynamic_spec = backend.args.mtp_dynamic_verify self.verifier = SpecVerifier(backend=backend) self.proposer = build_spec_proposer(engine=self) self.decode_runner = SpecDecodeRunner(engine=self) @@ -101,7 +100,7 @@ def prepare_draft_decode_input( model_input.input_ids = next_token_ids model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens - if self.spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3", "qwen3next_eagle"): + if self.spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3"): model_input.draft_step = 0 return model_input diff --git a/lightllm/server/router/model_infer/speculative/proposers/__init__.py b/lightllm/server/router/model_infer/speculative/proposers/__init__.py index cbab4cc3d1..d8ab194b3d 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/__init__.py +++ b/lightllm/server/router/model_infer/speculative/proposers/__init__.py @@ -18,7 +18,7 @@ def build_spec_proposer(engine) -> "BaseSpecProposer": from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer return Eagle3Proposer(engine=engine) - if spec_mode in ("eagle_with_att", "eagle_no_att", "qwen3next_eagle"): + if spec_mode in ("eagle_with_att", "eagle_no_att"): from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer return EagleMTPProposer(engine=engine) diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index 8b13ca863f..6aa0f20b7b 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -92,19 +92,36 @@ def propose_next( ) draft_model_output = draft_model.forward(draft_input) - flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + if self.enable_dynamic_spec: + flat_token_ids, flat_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) + else: + flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) assert flat_token_ids.numel() == num_reqs * block_size block_token_ids = flat_token_ids.reshape(num_reqs, block_size) # Standard DFlash has one leading bonus row; DeepSpec checkpoints do not. - bonus_rows = block_size - draft_step + bonus_rows = block_size - self.backend.max_draft_step assert bonus_rows in (0, 1), ( - f"DFlash block_size={block_size} must equal mtp_step={draft_step} " f"or mtp_step + 1={draft_step + 1}" + f"DFlash block_size={block_size} must equal mtp_step={self.backend.max_draft_step} " + f"or mtp_step + 1={self.backend.max_draft_step + 1}" ) - token_ids[selected_rows, 1:] = block_token_ids[:, bonus_rows:] + token_ids[selected_rows, 1:] = block_token_ids[:, bonus_rows : bonus_rows + draft_step] + + draft_probs = None + if self.enable_dynamic_spec: + block_token_probs = flat_token_probs.reshape(num_reqs, block_size) + selected_token_probs = block_token_probs[:, bonus_rows : bonus_rows + draft_step] + draft_probs = [ + self.scatter_selected_step_probs( + selected_rows=selected_rows, + selected_probs=selected_token_probs[:, step], + verify_row_count=next_token_ids.shape[0], + ) + for step in range(draft_step) + ] return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=draft_mem_indexes_cpu, - draft_probs=None, + draft_probs=draft_probs, ) def extend_draft_kv_cache(self, main_model_input: ModelInput, target_hidden: torch.Tensor) -> None: diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py index df380c83c5..0bb3bfe04b 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py @@ -55,6 +55,9 @@ def _get_pruned_active_count( active_count = math.ceil(self._draft_prune_safety_factor * max(1, draft_row_budget) / next_depth) return min(current_count, max(1, active_count)) + def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: + return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) + def propose_next( self, main_model_input: ModelInput, @@ -102,7 +105,7 @@ def propose_next( draft_logits = draft_model_output.logits.index_select(0, selected_rows) if collect_dynamic_probs: - draft_next_token_ids, selected_draft_prob = self.backend._gen_argmax_token_ids_and_prob( + draft_next_token_ids, selected_draft_prob = self._gen_argmax_token_ids_and_prob( ModelOutput(logits=draft_logits) ) draft_prob = self.scatter_selected_step_probs( @@ -113,7 +116,7 @@ def propose_next( draft_probs.append(draft_prob) chain_survival = selected_draft_prob.float().clamp(0.01, 0.99) else: - draft_next_token_ids = self.backend._gen_argmax_token_ids(ModelOutput(logits=draft_logits)) + draft_next_token_ids = self._gen_argmax_token_ids(ModelOutput(logits=draft_logits)) chain_survival = None assert draft_model_output.spec_hidden is not None draft_hidden = draft_model_output.spec_hidden.index_select(0, selected_rows) @@ -185,7 +188,7 @@ def propose_next( ) draft_output = draft_model.forward(draft_input) if collect_dynamic_probs: - draft_next_token_ids, selected_draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_output) + draft_next_token_ids, selected_draft_prob = self._gen_argmax_token_ids_and_prob(draft_output) draft_prob = self.scatter_selected_step_probs( selected_rows=selected_rows, selected_probs=selected_draft_prob, @@ -194,7 +197,7 @@ def propose_next( draft_probs.append(draft_prob) chain_survival = chain_survival * selected_draft_prob.float().clamp(0.01, 0.99) else: - draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_output) + draft_next_token_ids = self._gen_argmax_token_ids(draft_output) proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids draft_hidden = draft_output.spec_hidden assert draft_hidden is not None diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py index b93e7ca51d..c0a140e8ad 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -49,6 +49,16 @@ def build_initial_draft_state_overlap( def project_draft_decode_hidden(self, draft_hidden: torch.Tensor) -> torch.Tensor: return draft_hidden + def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: + return draft_token_ids + + def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: + return self._map_draft_token_ids(self.backend._gen_argmax_token_ids(model_output)) + + def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): + draft_token_ids, draft_probs = self.backend._gen_argmax_token_ids_and_prob(model_output) + return self._map_draft_token_ids(draft_token_ids), draft_probs + def make_verify_extend_input( self, base_input: ModelInput, @@ -228,7 +238,7 @@ def propose_next_overlap( zip(inputs, extend_outputs, selected_rows, real_request_nums) ): selected_output = ModelOutput(logits=extend_output.logits.index_select(0, selected)) - step_token_ids = self.backend._gen_argmax_token_ids(selected_output) + step_token_ids = self._gen_argmax_token_ids(selected_output) draft_next_token_ids.append(step_token_ids) assert extend_output.spec_hidden is not None draft_hiddens.append(extend_output.spec_hidden.index_select(0, selected)) @@ -293,7 +303,7 @@ def propose_next_overlap( step_outputs = draft_model.microbatch_overlap_decode(*step_inputs) for index, step_output in enumerate(step_outputs): - step_token_ids = self.backend._gen_argmax_token_ids(step_output) + step_token_ids = self._gen_argmax_token_ids(step_output) draft_next_token_ids[index] = step_token_ids draft_hiddens[index] = step_output.spec_hidden assert draft_hiddens[index] is not None @@ -377,7 +387,7 @@ def _propose_recurrent( selected_logits = extend_output.logits.index_select(0, selected_rows) selected_output = ModelOutput(logits=selected_logits) if self.enable_dynamic_spec: - draft_next_token_ids, selected_prob = self.backend._gen_argmax_token_ids_and_prob(selected_output) + draft_next_token_ids, selected_prob = self._gen_argmax_token_ids_and_prob(selected_output) draft_probs.append( self.scatter_selected_step_probs( selected_rows=selected_rows, @@ -386,7 +396,7 @@ def _propose_recurrent( ) ) else: - draft_next_token_ids = self.backend._gen_argmax_token_ids(selected_output) + draft_next_token_ids = self._gen_argmax_token_ids(selected_output) proposal_token_ids[selected_rows, 1] = draft_next_token_ids assert extend_output.spec_hidden is not None draft_hidden = extend_output.spec_hidden.index_select(0, selected_rows) @@ -426,7 +436,7 @@ def _propose_recurrent( ) draft_output = draft_model.forward(draft_input) if self.enable_dynamic_spec: - draft_next_token_ids, selected_prob = self.backend._gen_argmax_token_ids_and_prob(draft_output) + draft_next_token_ids, selected_prob = self._gen_argmax_token_ids_and_prob(draft_output) draft_probs.append( self.scatter_selected_step_probs( selected_rows=selected_rows, @@ -435,7 +445,7 @@ def _propose_recurrent( ) ) else: - draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_output) + draft_next_token_ids = self._gen_argmax_token_ids(draft_output) proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids draft_hidden = draft_output.spec_hidden assert draft_hidden is not None diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 30f8bea52e..ad461349b8 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -221,22 +221,6 @@ def enable_diverse_mode_gqa_decode_fast_kernel() -> bool: return get_env_start_args().diverse_mode and "int8kv" == get_env_start_args().llm_kv_type -@lru_cache(maxsize=None) -def enable_dynamic_spec() -> bool: - """Whether speculative scheduling may vary draft and verify widths. - - ``--mtp_dynamic_verify`` remains the compatible command-line switch; - DSpark enables dynamic speculative scheduling unconditionally. - """ - - args = get_env_start_args() - if args.mtp_mode == "dspark": - return True - if args.mtp_mode == "dflash": - return False - return bool(args.mtp_dynamic_verify) - - @lru_cache(maxsize=None) def enable_triton_mtp_kernel() -> bool: """ @@ -271,20 +255,23 @@ def enable_cpu_cache_numa_interleave() -> bool: @lru_cache(maxsize=None) def get_added_mtp_kv_layer_num() -> int: - # mtp 模式下需要在mem manger上扩展draft model使用的layer args = get_env_start_args() - spec_mode = args.mtp_mode - if spec_mode not in ("vanilla_with_att", "eagle_with_att", "eagle3", "dspark", "dflash"): + if args.mtp_mode == "eagle_with_att": + return 1 + if args.mtp_mode == "vanilla_with_att": + return args.mtp_step + if args.mtp_mode not in ("eagle3", "dspark", "dflash"): return 0 - draft_model_count = args.mtp_step if spec_mode == "vanilla_with_att" else 1 - if spec_mode in ("dflash", "dspark", "eagle3"): - if not args.mtp_draft_model_dir: - return draft_model_count - draft_model_dir = args.mtp_draft_model_dir[0] - with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: - draft_config = json.load(json_file) - return int(draft_config.get("num_hidden_layers", draft_config.get("n_layer", draft_model_count))) - return draft_model_count + + draft_model_dir = args.mtp_draft_model_dir[0] + with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: + draft_config = json.load(json_file) + # Use the effective draft backbone config when the checkpoint stores it nested. + draft_config.update(draft_config.get("dflash_config", {})) + # A draft model may contain multiple attention layers; each layer needs a + # separate KV-cache slot after the target model's layers. + total_added_mtp_kv_layer_num = draft_config.get("num_hidden_layers", draft_config.get("n_layer")) + return int(total_added_mtp_kv_layer_num) @lru_cache(maxsize=None) diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index 71dfe754c5..bca619dfe9 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -1,18 +1,19 @@ +from types import SimpleNamespace + +import pytest import torch from lightllm.common.basemodel import batch_objs -from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel from lightllm.models.qwen3_dspark import model_output as dspark_model_output -from lightllm.models.qwen3_dspark.model import DSparkModelOutputMixin +from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer +from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput -class _DSparkTestModel(DSparkModelOutputMixin, TpPartBaseModel): - pass - - -def test_dspark_decode_unpad_preserves_output_type_and_slices_dspark_fields(): - model = _DSparkTestModel.__new__(_DSparkTestModel) +@pytest.mark.parametrize("model_class", [Qwen3DSparkModel, Qwen3_5DSparkModel]) +def test_dspark_decode_unpad_preserves_output_type_and_slices_dspark_fields(model_class): + model = model_class.__new__(model_class) output = DSparkModelOutput( logits=torch.arange(48).view(12, 4), spec_hidden=torch.arange(36).view(12, 3), @@ -57,3 +58,59 @@ def test_dspark_no_ref_conversion_dispatches_to_dspark_fields(monkeypatch): output.draft_token_ids.data_ptr(), ) assert all(converted != original for converted, original in zip(converted_ptrs, original_ptrs)) + + +def test_vanilla_markov_local_sampling_matches_full_logits(): + post_infer = Qwen3DSparkPostLayerInfer.__new__(Qwen3DSparkPostLayerInfer) + post_infer.block_size_ = 3 + post_infer.markov_rank_ = 2 + post_infer.markov_head_type_ = "vanilla" + post_infer.tp_world_size_ = 1 + + vocab_size = 5 + request_count = 2 + local_logits = torch.randn(vocab_size, request_count * post_infer.block_size_) + + class MarkovEmbedding: + def __init__(self, weight): + self.weight = weight + + def __call__(self, input_ids, alloc_func): + return torch.nn.functional.embedding(input_ids, self.weight) + + class MarkovLMHead: + def __init__(self, weight): + self.weight = weight + self.tp_vocab_start_id = 0 + + def __call__(self, input, alloc_func): + return self.weight @ input + + markov_w1 = torch.randn(vocab_size, post_infer.markov_rank_) + markov_w2 = torch.randn(vocab_size, post_infer.markov_rank_) + layer_weight = SimpleNamespace( + markov_w1_weight_=MarkovEmbedding(markov_w1), + markov_w2_weight_=MarkovLMHead(markov_w2), + ) + anchor_token_ids = torch.tensor([1, 3]) + block_hidden = torch.empty(request_count, post_infer.block_size_, 0) + post_infer.alloc_tensor = torch.empty + + sampled_tokens = post_infer._sample_markov( + local_logits=local_logits, + block_hidden=block_hidden, + infer_state=SimpleNamespace(), + anchor_token_ids=anchor_token_ids, + layer_weight=layer_weight, + ) + + base_logits = local_logits.T.reshape(request_count, post_infer.block_size_, vocab_size) + prev_token_ids = anchor_token_ids + expected_tokens = [] + for step_idx in range(post_infer.block_size_): + markov_bias = torch.nn.functional.linear(markov_w1[prev_token_ids], markov_w2) + prev_token_ids = torch.argmax(base_logits[:, step_idx] + markov_bias, dim=-1) + expected_tokens.append(prev_token_ids) + expected_tokens = torch.stack(expected_tokens, dim=1) + + torch.testing.assert_close(sampled_tokens, expected_tokens) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index 37a8145e56..ad3d9bccfe 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -39,17 +39,16 @@ def test_padded_decode_builds_spec_metadata_for_real_and_fake_rows(monkeypatch): monkeypatch.setattr( generic_padded_pre_process, "get_env_start_args", - lambda: SimpleNamespace(mtp_step=max_draft_step), + lambda: SimpleNamespace(mtp_step=max_draft_step, mtp_dynamic_verify=True), ) monkeypatch.setattr(generic_padded_pre_process, "enable_diverse_mode_gqa_decode_fast_kernel", lambda: False) - monkeypatch.setattr(generic_padded_pre_process, "enable_dynamic_spec", lambda: True) monkeypatch.setattr(generic_padded_pre_process, "enable_triton_mtp_kernel", lambda: False) monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 8) req = SimpleNamespace( req_idx=7, cur_kv_len=4, - max_draft_step=max_draft_step, + mtp_step=max_draft_step, multimodal_params={"images": [], "audios": []}, get_cur_total_len=lambda: 5, ) diff --git a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py index 57501d3a19..da082ad018 100644 --- a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py @@ -3,6 +3,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer @@ -86,3 +87,19 @@ def test_overlap_eagle_extends_verify_rows_then_decodes_logical_batch(monkeypatc assert proposal.token_ids.shape == (9, 3) assert torch.equal(proposal.token_ids[:, 0], torch.tensor([0, 1, 2, 10, 11, 12, 13, 14, 15])) assert proposal.extra_mem_indexes_cpu.shape == (3,) + + +def test_eagle3_maps_draft_token_ids_in_proposer(): + proposer = Eagle3Proposer.__new__(Eagle3Proposer) + proposer.backend = SimpleNamespace( + draft_models=[SimpleNamespace(map_draft_vocab_to_main_vocab=lambda token_ids: token_ids + 100)], + _gen_argmax_token_ids=lambda _: torch.tensor([1, 2]), + _gen_argmax_token_ids_and_prob=lambda _: (torch.tensor([3, 4]), torch.tensor([0.8, 0.7])), + ) + + token_ids = proposer._gen_argmax_token_ids(ModelOutput(logits=torch.empty(0))) + token_ids_with_prob, probs = proposer._gen_argmax_token_ids_and_prob(ModelOutput(logits=torch.empty(0))) + + assert torch.equal(token_ids, torch.tensor([101, 102])) + assert torch.equal(token_ids_with_prob, torch.tensor([103, 104])) + assert torch.equal(probs, torch.tensor([0.8, 0.7])) diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 3ee3f36e7f..0306487c10 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -3,13 +3,15 @@ from types import SimpleNamespace import pytest +import torch import lightllm.common.basemodel.attention.base_att as base_att_module -import lightllm.server.router.model_infer.mode_backend.base_backend as base_backend_module +import lightllm.common.basemodel.hidden_collector as hidden_collector_module from lightllm.common.basemodel.attention.base_att import BaseAttBackend +from lightllm.common.basemodel.hidden_collector import HiddenCollector from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight -from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend +from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer from lightllm.utils import envs_utils @@ -49,6 +51,8 @@ def test_qwen3_eagle_uses_layers_checkpoint_prefix(): [ ("dspark", False, 7, True, True), ("dspark", True, 7, True, False), + ("dflash", False, 7, True, True), + ("dflash", True, 7, True, False), ("vanilla_with_att", True, 7, True, True), ("eagle3", True, 0, True, False), ("eagle3", False, 7, False, False), @@ -62,112 +66,142 @@ def test_attention_backend_selects_dynamic_spec_layout( dynamic_spec, expected, ): - monkeypatch.setattr(base_att_module, "enable_dynamic_spec", lambda: dynamic_spec) - monkeypatch.setattr(base_att_module, "get_env_start_args", lambda: SimpleNamespace(mtp_mode=mtp_mode)) - backend = SimpleNamespace(model=SimpleNamespace(is_mtp_draft_model=is_draft_model)) + monkeypatch.setattr( + base_att_module, + "get_env_start_args", + lambda: SimpleNamespace(mtp_mode=mtp_mode, mtp_dynamic_verify=dynamic_spec), + ) + backend = SimpleNamespace( + model=SimpleNamespace( + is_mtp_draft_model=is_draft_model, + ) + ) infer_state = SimpleNamespace(draft_step=draft_step) assert BaseAttBackend.uses_dynamic_spec_verify_layout(backend, infer_state) is expected @pytest.mark.parametrize( - "spec_mode, architecture, is_linear_att_mixed_model, expected_class_name", + "model_type, spec_mode, expected_class_name", [ - ("dflash", "Qwen3DFlashModel", False, "Qwen3DFlashModel"), - ("dflash", "Qwen3DSparkModel", False, "Qwen3DFlashModel"), - ("dflash", "Qwen3_5DFlashModel", True, "Qwen3_5DFlashModel"), - ("dspark", "Qwen3DSparkModel", False, "Qwen3DSparkModel"), - ("dspark", "Qwen3DSparkModel", True, "Qwen3_5DSparkModel"), - ("eagle3", "Qwen3Eagle3Model", False, "Qwen3EagleModel"), + ("deepseek_v3", "vanilla_with_att", "Deepseek3MTPModel"), + ("deepseek_v3", "eagle_with_att", "Deepseek3MTPModel"), + ("glm4_moe_lite", "vanilla_with_att", "Glm4MoeLiteMTPModel"), + ("glm4_moe_lite", "eagle_with_att", "Glm4MoeLiteMTPModel"), + ("mistral", "vanilla_no_att", "MistralMTPModel"), + ("mistral", "eagle_no_att", "MistralMTPModel"), + ("qwen3_moe", "vanilla_no_att", "Qwen3MOEMTPModel"), + ("qwen3_moe", "eagle_no_att", "Qwen3MOEMTPModel"), + ("qwen3_5", "vanilla_with_att", "Qwen3_5MTPModel"), + ("qwen3_5_text", "eagle_with_att", "Qwen3_5MTPModel"), + ("qwen3_5_moe", "vanilla_with_att", "Qwen3_5MoeMTPModel"), + ("qwen3_5_moe_text", "eagle_with_att", "Qwen3_5MoeMTPModel"), + ("qwen3", "dflash", "Qwen3DFlashModel"), + ("qwen3_5", "dflash", "Qwen3_5DFlashModel"), + ("qwen3_5_text", "dflash", "Qwen3_5DFlashModel"), + ("qwen3", "dspark", "Qwen3DSparkModel"), + ("qwen3_5", "dspark", "Qwen3_5DSparkModel"), + ("qwen3_5_text", "dspark", "Qwen3_5DSparkModel"), + ("qwen3", "eagle3", "Qwen3EagleModel"), ], ) -def test_draft_model_registry( - spec_mode, - architecture, - is_linear_att_mixed_model, - expected_class_name, -): +def test_draft_model_registry(model_type, spec_mode, expected_class_name): model_class = get_draft_model_class( - model_cfg={"model_type": "qwen3", "architectures": [architecture]}, + model_cfg={"model_type": model_type}, spec_mode=spec_mode, - is_linear_att_mixed_model=is_linear_att_mixed_model, ) assert model_class.__name__ == expected_class_name -def test_dspark_env_helper_enables_dynamic_spec_without_legacy_flag(monkeypatch): - monkeypatch.setenv( - "LIGHTLLM_START_ARGS", - json.dumps( - { - "mtp_mode": "dspark", - "mtp_dynamic_verify": False, - "mtp_step": 4, - } - ), - ) - envs_utils.get_env_start_args.cache_clear() - envs_utils.enable_dynamic_spec.cache_clear() - - assert envs_utils.enable_dynamic_spec() +def test_draft_model_registry_rejects_unsupported_model_type(): + with pytest.raises(ValueError, match="Unsupported speculative draft model"): + get_draft_model_class( + model_cfg={"model_type": "gemma4"}, + spec_mode="dspark", + ) @pytest.mark.parametrize( - "mtp_mode, mtp_dynamic_verify, expected", + "model_type, spec_mode", [ - ("dspark", False, True), - ("dflash", True, False), - ("eagle3", True, True), - ("eagle3", False, False), + ("deepseek_v3", "eagle3"), + ("deepseek_v3", "dspark"), + ("deepseek_v3", "dflash"), + ("glm4_moe_lite", "eagle3"), + ("glm4_moe_lite", "dspark"), + ("glm4_moe_lite", "dflash"), + ("qwen3_5", "eagle3"), + ("qwen3_5_moe", "eagle3"), + ("qwen3_5_moe", "dspark"), + ("qwen3_5_moe", "dflash"), + ("qwen3", "eagle_no_att"), ], ) -def test_env_helper_uses_normalized_dynamic_spec(monkeypatch, mtp_mode, mtp_dynamic_verify, expected): - monkeypatch.setenv( - "LIGHTLLM_START_ARGS", - json.dumps( - { - "mtp_mode": mtp_mode, - "mtp_step": 7, - "mtp_dynamic_verify": mtp_dynamic_verify, - } - ), - ) - envs_utils.get_env_start_args.cache_clear() - envs_utils.enable_dynamic_spec.cache_clear() - - assert envs_utils.enable_dynamic_spec() is expected - - -def test_draft_model_registry_rejects_unsupported_architecture(): +def test_draft_model_registry_rejects_unsupported_mode(model_type, spec_mode): with pytest.raises(ValueError, match="Unsupported speculative draft model"): get_draft_model_class( - model_cfg={"model_type": "gemma4", "architectures": ["Gemma4DSparkModel"]}, - spec_mode="dspark", + model_cfg={"model_type": model_type}, + spec_mode=spec_mode, ) -def test_draft_model_registry_validates_target_family(): - with pytest.raises(ValueError, match="linear-attention mixed targets"): - get_draft_model_class( - model_cfg={"model_type": "qwen3", "architectures": ["Qwen3DFlashModel"]}, - spec_mode="dflash", - is_linear_att_mixed_model=True, - ) +def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): + block_size = 4 + max_draft_step = 3 + verify_row_count = 5 + selected_rows = torch.tensor([0, 3]) + flat_token_ids = torch.arange(2 * block_size) + flat_token_probs = torch.arange(2 * block_size, dtype=torch.float32) / 10 + + draft_model = SimpleNamespace( + block_size=block_size, + forward=lambda _: SimpleNamespace(logits=torch.empty(2 * block_size, 1)), + ) + backend = SimpleNamespace( + max_draft_step=max_draft_step, + draft_models=[draft_model], + _gen_argmax_token_ids_and_prob=lambda _: (flat_token_ids, flat_token_probs), + ) + proposer = DFlashProposer(SimpleNamespace(backend=backend, enable_dynamic_spec=True)) + proposer.extend_draft_kv_cache = lambda **_: None + proposer.select_accepted_tail_rows = lambda **_: selected_rows + proposer.build_block_draft_input = lambda **_: (SimpleNamespace(), torch.tensor([10, 11])) + + proposal = proposer.propose_next( + main_model_input=SimpleNamespace(), + main_model_output=SimpleNamespace(spec_hidden=torch.empty(verify_row_count, 1)), + next_token_ids=torch.arange(verify_row_count), + b_req_mtp_start_loc=torch.tensor([0, 3]), + draft_step=2, + accept_len=torch.tensor([1, 1]), + ) + + expected_blocks = flat_token_ids.reshape(2, block_size)[:, 1:3] + torch.testing.assert_close(proposal.token_ids[selected_rows, 1:], expected_blocks) + assert len(proposal.draft_probs) == 2 + for step, probs in enumerate(proposal.draft_probs): + expected_probs = torch.zeros(verify_row_count) + expected_probs[selected_rows] = flat_token_probs.reshape(2, block_size)[:, step + 1] + torch.testing.assert_close(probs, expected_probs) -def test_target_hidden_layer_ids_reads_text_config(monkeypatch): +def test_hidden_collector_reads_target_layer_ids(monkeypatch): monkeypatch.setattr( - base_backend_module.PretrainedConfig, + hidden_collector_module.PretrainedConfig, "get_config_dict", lambda _: ({"target_layer_ids": [1, 20, 36]}, {}), ) - backend = ModeBackend.__new__(ModeBackend) - backend.args = SimpleNamespace(mtp_mode="dspark", mtp_draft_model_dir=["/models/dspark"]) + monkeypatch.setattr( + hidden_collector_module, + "get_env_start_args", + lambda: SimpleNamespace(mtp_draft_model_dir=["/models/dspark"]), + ) + model = SimpleNamespace(is_mtp_draft_model=False, layers_num=40) - layer_ids = backend._target_hidden_layer_ids({"model_type": "qwen3_5", "text_config": {"num_hidden_layers": 40}}) + collector = HiddenCollector(model=model, spec_mode="dspark") - assert layer_ids == (1, 20, 36) + assert collector.collectors[0].layer_ids == frozenset((1, 20, 36)) def test_dflash_added_kv_layers_come_from_draft_config(tmp_path): @@ -188,6 +222,31 @@ def test_dflash_added_kv_layers_come_from_draft_config(tmp_path): assert envs_utils.get_added_mtp_kv_layer_num() == 5 +def test_qwen35_dflash_added_kv_layers_come_from_nested_draft_config(tmp_path): + config_path = tmp_path / "config.json" + config_path.write_text( + json.dumps( + { + "num_hidden_layers": 48, + "dflash_config": {"num_hidden_layers": 5}, + } + ) + ) + + envs_utils.get_env_start_args.cache_clear() + envs_utils.get_added_mtp_kv_layer_num.cache_clear() + envs_utils.set_env_start_args( + { + "mtp_mode": "dflash", + "mtp_step": 7, + "mtp_dynamic_verify": False, + "mtp_draft_model_dir": [str(tmp_path)], + } + ) + + assert envs_utils.get_added_mtp_kv_layer_num() == 5 + + def test_dspark_added_kv_layers_come_from_draft_config(tmp_path): config_path = tmp_path / "config.json" config_path.write_text(json.dumps({"num_hidden_layers": 6})) From 9aa6dfc425fedcdef6eb3b3c19717f5c50fc44d4 Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Tue, 11 Aug 2026 10:43:39 +0000 Subject: [PATCH 007/103] refactor: simplify speculative decoding pipeline --- .../pre_and_post_layer_weight.py | 4 +- .../model_infer/mode_backend/base_backend.py | 6 +- .../mode_backend/chunked_prefill/impl.py | 1 - .../router/model_infer/speculative/engine.py | 80 +++++++------- .../router/model_infer/speculative/planner.py | 70 +++++------- .../speculative/proposers/__init__.py | 6 +- .../model_infer/speculative/proposers/base.py | 6 +- .../speculative/proposers/dflash.py | 9 +- .../speculative/proposers/dspark.py | 29 +---- .../speculative/proposers/eagle_mtp.py | 30 +----- .../router/model_infer/speculative/runner.py | 12 +-- .../model_infer/speculative/verifier.py | 100 ------------------ .../models/test_qwen3_dspark_model_output.py | 34 ++++++ 13 files changed, 133 insertions(+), 254 deletions(-) delete mode 100644 lightllm/server/router/model_infer/speculative/verifier.py diff --git a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py index 3965a8163a..e6c23b0198 100644 --- a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py +++ b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py @@ -74,7 +74,9 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): weight_names="confidence_head.proj.weight", bias_names="confidence_head.proj.bias", data_type=self.data_type_, - quant_method=self.quant_cfg.get_quant_method(0, "confidence_head.proj"), + # The confidence head has a bias and only one output channel. + # Keep it in model dtype because W8A8 MM does not support bias. + quant_method=None, tp_rank=0, tp_world_size=1, ) 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 64ba873d81..6e46b2e75b 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -303,7 +303,11 @@ def init_spec_engine(self, main_kvargs: dict): if self.args.mtp_mode is None: return self.init_mtp_draft_model(main_kvargs) - self.spec_engine = SpecEngine(backend=self) + self.spec_engine = SpecEngine( + backend=self, + spec_mode=self.args.mtp_mode, + enable_dynamic_spec=self.args.mtp_dynamic_verify, + ) return def init_mtp_draft_model(self, main_kvargs: dict): diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 3f8c3c48d5..8e42f73013 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -300,7 +300,6 @@ def decode_mtp( spec_post_state = spec_engine.finish_decode_post( state=spec_decode_state, req_num=len(decode_reqs), - run_reqs=run_reqs, ) self._update_spec_accept_ratio( decode_reqs=decode_reqs, diff --git a/lightllm/server/router/model_infer/speculative/engine.py b/lightllm/server/router/model_infer/speculative/engine.py index f7dcd3c0f3..2f8cc05cbc 100644 --- a/lightllm/server/router/model_infer/speculative/engine.py +++ b/lightllm/server/router/model_infer/speculative/engine.py @@ -5,6 +5,10 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.mtp_utils import ( + mtp_scatter_next_token_ids, + mtp_verify, +) from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.server.router.model_infer.speculative.planner import ( DSparkDynamicSpecPlanner, @@ -20,29 +24,25 @@ SpecDecodePostState, SpecDecodeRunner, ) -from lightllm.server.router.model_infer.speculative.verifier import SpecVerifier, SpecVerifyResult class SpecEngine: - """Owns one speculative decoding pipeline for a model backend. - - BaseModel only returns ``ModelOutput.spec_hidden``. The engine owns - planning, target verification, draft extension/decoding, and the request - bookkeeping around that pipeline. Algorithm-specific state stays in the - proposer. - - The target-to-draft data path is explicit: - 1. target forward returns ``ModelOutput.spec_hidden`` - 2. proposer performs one draft extend from that hidden state - 3. recurrent proposers perform zero or more unit-stride draft decodes - 4. verifier accepts target rows and scatters the next candidate block + """Coordinates speculative decoding for one model backend. + + The engine handles the common decode flow, while each proposer owns its + algorithm-specific draft-state update and proposal generation. + + Each decode iteration: + 1. the target model samples tokens and returns captured hidden states; + 2. verification computes the accepted prefix against the previous proposal; + 3. the proposer updates draft state and generates the next proposal; + 4. the proposal is stored in per-request buffers for the next iteration. """ - def __init__(self, backend) -> None: + def __init__(self, backend, spec_mode: str, enable_dynamic_spec: bool) -> None: self.backend = backend - self.spec_mode = backend.args.mtp_mode - self.enable_dynamic_spec = backend.args.mtp_dynamic_verify - self.verifier = SpecVerifier(backend=backend) + self.spec_mode = spec_mode + self.enable_dynamic_spec = enable_dynamic_spec self.proposer = build_spec_proposer(engine=self) self.decode_runner = SpecDecodeRunner(engine=self) self.planner = self._build_decode_planner() @@ -175,9 +175,8 @@ def finish_decode_post( self, state: SpecDecodeForwardState, req_num: int, - run_reqs: List, ) -> SpecDecodePostState: - return self.decode_runner.finish_post(state=state, req_num=req_num, run_reqs=run_reqs) + return self.decode_runner.finish_post(state=state, req_num=req_num) def prepare_decode_model_input( self, @@ -222,14 +221,12 @@ def build_decode_req_lists( ): """Build post-handle request lists after optional verify-row compaction.""" - if self.enable_dynamic_spec and selected_row_mask_cpu is not None: - selected_row_mask_numpy = selected_row_mask_cpu.numpy() - run_reqs = [original_run_reqs[i] for i in range(len(original_run_reqs)) if selected_row_mask_numpy[i] == 1] + if selected_row_mask_cpu is not None: + run_reqs = [req for req, selected in zip(original_run_reqs, selected_row_mask_cpu.tolist()) if selected] else: run_reqs = original_run_reqs - accepted_index_cpu_numpy = accepted_index_cpu.numpy() - verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu_numpy[i] == 1] + verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] return run_reqs, verify_ok_reqs def build_decode_free_mem_indexes_cpu( @@ -239,7 +236,7 @@ def build_decode_free_mem_indexes_cpu( accepted_index_cpu: torch.Tensor, ) -> torch.Tensor: mem_indexes_cpu = model_input.mem_indexes_cpu - if not self.enable_dynamic_spec or selected_row_mask_cpu is None: + if selected_row_mask_cpu is None: return mem_indexes_cpu[accepted_index_cpu == 0] selected_mask = selected_row_mask_cpu.to(dtype=torch.bool) @@ -258,17 +255,12 @@ def build_decode_free_mem_indexes_cpu( def update_dynamic_accept_stats( self, req_num: int, - run_reqs, accepted_index_cpu: torch.Tensor, spec_accept_len_cpu: torch.Tensor, - dynamic_batch_size: Optional[int], - pre_draft_step: Optional[int] = None, + dynamic_batch_size: int, + pre_draft_step: int, ) -> None: - if not self.enable_dynamic_spec: - return - - assert dynamic_batch_size is not None - assert len(run_reqs) == accepted_index_cpu.shape[0] + assert accepted_index_cpu.shape[0] == dynamic_batch_size assert spec_accept_len_cpu.shape[0] == req_num accept_lengths = spec_accept_len_cpu.numpy() accept_count = int(accept_lengths.sum()) @@ -331,11 +323,8 @@ def needs_schedule_probs_cpu(self) -> bool: def update_dynamic_schedule_stats( self, req_num: int, - schedule_probs_cpu: Optional[torch.Tensor], + schedule_probs_cpu: torch.Tensor, ) -> None: - if not self.needs_schedule_probs_cpu() or schedule_probs_cpu is None: - return - self.planner.update_predicted_schedule_probs( schedule_probs=schedule_probs_cpu, req_num=req_num, @@ -344,7 +333,7 @@ def update_dynamic_schedule_stats( def propose_next( self, main_model_input: ModelInput, - main_model_output: Optional[ModelOutput], + main_model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, @@ -392,11 +381,12 @@ def verify_target_tokens( new_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, - ) -> SpecVerifyResult: - return self.verifier.verify_target_tokens( + ) -> Tuple[torch.Tensor, torch.Tensor]: + return mtp_verify( + req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, new_next_token_ids=new_next_token_ids, b_req_idx=b_req_idx, - b_req_mtp_start_loc=b_req_mtp_start_loc, ) def build_all_next_token_probs( @@ -458,11 +448,17 @@ def scatter_next_tokens( spec_accept_len: torch.Tensor, all_next_token_probs: Optional[torch.Tensor] = None, ) -> None: - self.verifier.scatter_next_tokens( + mtp_scatter_next_token_ids( + req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=all_next_token_ids, b_req_idx=b_req_idx, spec_accept_len=spec_accept_len, + req_to_next_token_probs=( + self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs + if all_next_token_probs is not None + else None + ), all_next_token_probs=all_next_token_probs, ) diff --git a/lightllm/server/router/model_infer/speculative/planner.py b/lightllm/server/router/model_infer/speculative/planner.py index 5781831e96..33bb5f737d 100644 --- a/lightllm/server/router/model_infer/speculative/planner.py +++ b/lightllm/server/router/model_infer/speculative/planner.py @@ -63,8 +63,8 @@ def __init__( self.max_draft_step = int(max_draft_step) # 用于记录 decode 时的静态推理耗时(ms)。 - self.main_model_speeds = _InferCostMsTable() - self.draft_model_speeds = _InferCostMsTable() + self.target_infer_costs = _InferCostMsTable() + self.draft_infer_costs = _InferCostMsTable() # 记录不同 draft 深度的接受概率;第一个 target token 必然接受,不需要统计。 self.draft_len_to_accept_ratio = [ @@ -94,8 +94,8 @@ def plan(self, req_num: int, original_batch_size: int) -> SpecDecodePlan: ) def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: - speed_table = self.draft_model_speeds if is_draft_model else self.main_model_speeds - speed_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs + cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) def update_draft_len_to_accept_ratio(self, draft_len: int, accept_ratio: float) -> None: assert draft_len > 0 and draft_len <= self.max_draft_step @@ -143,7 +143,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl if req_num == 0: self.pre_draft_step = self.max_draft_step return 0, self.max_draft_step, pre_draft_step - if not self.main_model_speeds.has_data() or not self.draft_model_speeds.has_data(): + if not self.target_infer_costs.has_data() or not self.draft_infer_costs.has_data(): # The cost model is only meaningful after both target and draft # decode costs have been profiled. Block proposers such as DFlash # do not run through draft_model.forward, and cudagraph may also be @@ -152,7 +152,6 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl self.pre_draft_step = self.max_draft_step return req_num * (pre_draft_step + 1), self.max_draft_step, pre_draft_step - # case 1 如果采用随机的方式决定 dynamic_batch_size self._iter += 1 if self._use_random_mode and self._iter % self._iter_threshold == 0: min_batch_size = req_num @@ -163,30 +162,24 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl self.pre_draft_step = draft_step return dynamic_batch_size, draft_step, pre_draft_step - # 通过计算的方式来获取已经知道的最优的 dynamic_batch_size,然后再决定 draft step 步长 min_batch_size = req_num max_batch_size = req_num * (pre_draft_step + 1) - dynamic_batch_size_keys = self.main_model_speeds.get_batch_size_keys_between(min_batch_size, max_batch_size) + dynamic_batch_size_keys = self.target_infer_costs.get_batch_size_keys_between(min_batch_size, max_batch_size) - # 计算每个 dynamic_batch_size 对应的接受率以及单token的速度收益,然后选择最优的 dynamic_batch_size cost_ms_list = [ self._get_cost_ms(req_num=req_num, dynamic_batch_size=dynamic_batch_size, draft_step=pre_draft_step) for dynamic_batch_size in dynamic_batch_size_keys ] dynamic_batch_size = dynamic_batch_size_keys[np.argmin(cost_ms_list)] - # 下一步的 draft step 选择,需要考虑计算不同step步的收益问题再决定 min_cost_ms = float("inf") - min_cost_ms_draft_step = 0 # 默认选择0步长 + min_cost_ms_draft_step = 0 for draft_step in range(0, self.max_draft_step + 1): cost_ms = self._get_cost_ms(req_num=req_num, dynamic_batch_size=dynamic_batch_size, draft_step=draft_step) if cost_ms < min_cost_ms: min_cost_ms = cost_ms min_cost_ms_draft_step = draft_step - # draft step 步长不能超过 max_draft_step, 也不能小于0 - min_cost_ms_draft_step = min(min_cost_ms_draft_step, self.max_draft_step) - min_cost_ms_draft_step = max(min_cost_ms_draft_step, 0) self.pre_draft_step = min_cost_ms_draft_step return dynamic_batch_size, min_cost_ms_draft_step, pre_draft_step @@ -195,8 +188,8 @@ def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) - req_num=req_num, dynamic_batch_size=dynamic_batch_size ) total_time = ( - self.main_model_speeds.get(dynamic_batch_size) - + self.draft_model_speeds.get(dynamic_batch_size) * draft_step + self.target_infer_costs.get(dynamic_batch_size) + + self.draft_infer_costs.get(dynamic_batch_size) * draft_step ) token_num = min((dynamic_batch_size * accept_ratio), req_num * (draft_step + 1)) token_num = max(token_num, req_num) @@ -230,10 +223,10 @@ def _get_dynamic_batch_size_to_accept_ratio(self, req_num: int, dynamic_batch_si right_value = self.draft_len_to_accept_ratio[right - 1].get() accept_ratio = left_value + (right_value - left_value) * (real_step - left) - calcu_accept_ratio = (req_num + (dynamic_batch_size - req_num) * accept_ratio) / dynamic_batch_size + estimated_accept_ratio = (req_num + (dynamic_batch_size - req_num) * accept_ratio) / dynamic_batch_size weight = ema.get_count() / 10 # 通过统计数据和单请求数据进行加权平均,得到最终的接受率 - return calcu_accept_ratio * (1 - weight) + ema.get() * weight + return estimated_accept_ratio * (1 - weight) + ema.get() * weight class Eagle3DynamicSpecPlanner(DynamicSpecPlanner): @@ -557,7 +550,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl return 0, self.max_draft_step, pre_draft_step max_batch_size = req_num * (pre_draft_step + 1) - if not self.main_model_speeds.has_data() or not self.draft_model_speeds.has_data(): + if not self.target_infer_costs.has_data() or not self.draft_infer_costs.has_data(): self.pre_draft_step = self.max_draft_step return max_batch_size, self.max_draft_step, pre_draft_step @@ -595,7 +588,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl # captured shape. Turn those paid-for padding rows into real # candidates so they can improve accepted progress for the # same target-model graph cost. - graph_batch_size = self.main_model_speeds.get_ceil_batch_size( + graph_batch_size = self.target_infer_costs.get_ceil_batch_size( dynamic_batch_size, max_batch_size=max_batch_size, ) @@ -651,7 +644,7 @@ def _get_controlled_verify_rows_per_req(self) -> float: def _get_candidate_batch_sizes(self, req_num: int, max_batch_size: int) -> List[int]: candidates = set( - self.main_model_speeds.get_batch_size_keys_between( + self.target_infer_costs.get_batch_size_keys_between( req_num, max_batch_size, ) @@ -866,11 +859,11 @@ def _get_eagle3_cost_ms( # proposal step after that runs on one accepted tail per request. draft_cost_ms = 0.0 if draft_step > 0: - draft_cost_ms = self.draft_model_speeds.get(dynamic_batch_size) + draft_cost_ms = self.draft_infer_costs.get(dynamic_batch_size) if draft_step > 1: - draft_cost_ms += self.draft_model_speeds.get(req_num) * (draft_step - 1) + draft_cost_ms += self.draft_infer_costs.get(req_num) * (draft_step - 1) - total_time_ms = self.main_model_speeds.get(dynamic_batch_size) + draft_cost_ms + total_time_ms = self.target_infer_costs.get(dynamic_batch_size) + draft_cost_ms return total_time_ms / expected_token_num @@ -908,7 +901,7 @@ def update_predicted_schedule_probs(self, schedule_probs, req_num: int) -> None: if req_num <= 0: return - if not self.main_model_speeds.has_data(): + if not self.target_infer_costs.has_data(): return probs = self._to_numpy(schedule_probs) @@ -939,7 +932,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl return 0, self.max_draft_step, pre_draft_step max_batch_size = req_num * (pre_draft_step + 1) - if not self.main_model_speeds.has_data(): + if not self.target_infer_costs.has_data(): return max_batch_size, self.max_draft_step, pre_draft_step historical_batch_size = self._pop_historical_dynamic_batch_size( @@ -954,7 +947,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl # leaking a same-step EMA fallback into DSpark scheduling. return req_num, self.max_draft_step, pre_draft_step - candidate_batch_sizes = set(self.main_model_speeds.get_batch_size_keys_between(req_num, max_batch_size)) + candidate_batch_sizes = set(self.target_infer_costs.get_batch_size_keys_between(req_num, max_batch_size)) candidate_batch_sizes.add(req_num) candidate_batch_sizes.add(max_batch_size) survival_prefix = self._estimate_survival_prefix(pre_draft_step) @@ -968,7 +961,7 @@ def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tupl dynamic_batch_size=dynamic_batch_size, survival_prefix=survival_prefix, ) - verify_ms = max(self.main_model_speeds.get(dynamic_batch_size), 1e-6) + verify_ms = max(self.target_infer_costs.get(dynamic_batch_size), 1e-6) throughput = expected_tokens / verify_ms if throughput > best_throughput: best_throughput = throughput @@ -992,7 +985,7 @@ def _select_dynamic_batch_size_from_survival_scores( if flat_survival_scores.shape[0] == 0: return int(req_num) - candidate_batch_sizes = set(self.main_model_speeds.get_batch_size_keys_between(req_num, max_batch_size)) + candidate_batch_sizes = set(self.target_infer_costs.get_batch_size_keys_between(req_num, max_batch_size)) candidate_batch_sizes.add(int(req_num)) candidate_batch_sizes.add(max_batch_size) @@ -1013,7 +1006,7 @@ def _select_dynamic_batch_size_from_survival_scores( for dynamic_batch_size in sorted(candidate_batch_sizes): selected_draft_count = dynamic_batch_size - int(req_num) expected_tokens = float(req_num) + float(expected_accepts_by_count[selected_draft_count]) - verify_ms = max(self.main_model_speeds.get(dynamic_batch_size), 1e-6) + verify_ms = max(self.target_infer_costs.get(dynamic_batch_size), 1e-6) throughput = expected_tokens / verify_ms if throughput > best_throughput: best_throughput = throughput @@ -1109,34 +1102,27 @@ def get(self, batch_size: int) -> float: assert batch_size > 0 batch_size = int(batch_size) - # 这种情况理论上不应该存在。 if len(self.infer_cost_ms_table) == 0: return batch_size * 1000.0 - # 存在这个 batch_size 的记录,直接返回 if batch_size in self.infer_cost_ms_table: return self.infer_cost_ms_table[batch_size] - # 不存在这个 batch_size,需要进行插值估计。 max_batch_size = self.infer_cost_ms_table.peekitem(-1)[0] if batch_size > max_batch_size: max_infer_cost_ms = self.infer_cost_ms_table.peekitem(-1)[1] # 超过最大 graph 范围时使用高惩罚,使调度器倾向于关闭 speculative draft。 return max_infer_cost_ms + (batch_size - max_batch_size) * 1000.0 - else: - # 找到第一个大于等于 batch_size 的 key,并返回它的 value。 - index = self.infer_cost_ms_table.bisect_left(batch_size) - return self.infer_cost_ms_table.peekitem(index)[1] + + index = self.infer_cost_ms_table.bisect_left(batch_size) + return self.infer_cost_ms_table.peekitem(index)[1] def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> List[int]: assert batch_size1 > 0 and batch_size2 > 0 start = min(int(batch_size1), int(batch_size2)) end = max(int(batch_size1), int(batch_size2)) - ans = list(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) - if len(ans) == 0: - return [end] - else: - return ans + batch_sizes = list(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) + return batch_sizes or [end] def get_ceil_batch_size(self, batch_size: int, max_batch_size: int) -> Optional[int]: """Return the next recorded graph shape without inventing a key. diff --git a/lightllm/server/router/model_infer/speculative/proposers/__init__.py b/lightllm/server/router/model_infer/speculative/proposers/__init__.py index d8ab194b3d..85a38a81ca 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/__init__.py +++ b/lightllm/server/router/model_infer/speculative/proposers/__init__.py @@ -22,10 +22,12 @@ def build_spec_proposer(engine) -> "BaseSpecProposer": from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer return EagleMTPProposer(engine=engine) + if spec_mode in ("vanilla_with_att", "vanilla_no_att"): + from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer - from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer + return VanillaMTPProposer(engine=engine) - return VanillaMTPProposer(engine=engine) + raise ValueError(f"unsupported speculative mode: {spec_mode}") __all__ = [ diff --git a/lightllm/server/router/model_infer/speculative/proposers/base.py b/lightllm/server/router/model_infer/speculative/proposers/base.py index fcd45be8f7..2a0b565556 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/speculative/proposers/base.py @@ -57,7 +57,7 @@ class BaseSpecProposer: A proposer owns the draft-side state transition. The target model gives it the current target token ids plus captured target hidden features through SpecEngine.prepare_draft_* methods. The proposer returns candidate ids - but does not verify acceptance; verification is handled by SpecVerifier. + but does not verify acceptance; verification is handled by SpecEngine. """ def __init__(self, engine: "SpecEngine") -> None: @@ -102,7 +102,7 @@ def scatter_selected_step_probs( verify_row_count: int, ) -> torch.Tensor: out = torch.zeros( - (verify_row_count,), + (verify_row_count, *selected_probs.shape[1:]), dtype=torch.float32, device=selected_probs.device, ) @@ -154,7 +154,7 @@ def build_initial_draft_state_overlap( def propose_next( self, main_model_input: ModelInput, - main_model_output: Optional[ModelOutput], + main_model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index 6aa0f20b7b..e03035021d 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -9,7 +9,7 @@ class DFlashProposer(BaseSpecProposer): - """DFlash block proposer aligned to the Eagle3 engine boundary. + """Non-causal block proposer for DFlash. DFlash remains a non-causal block-prefill draft model, not a recurrent token decoder. The service flow is: @@ -75,7 +75,6 @@ def propose_next( return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=None, - draft_probs=None, ) # DFlash drafts from the accepted tail row of each request; one anchor @@ -132,7 +131,8 @@ def extend_draft_kv_cache(self, main_model_input: ModelInput, target_hidden: tor draft_kv_input = copy.copy(main_model_input) draft_kv_input.batch_size = batch_size draft_kv_input.total_token_num = batch_size - draft_kv_input.multimodal_params = [{"images": [], "audios": []} for _ in range(batch_size)] + empty_multimodal_params = {"images": [], "audios": []} + draft_kv_input.multimodal_params = [empty_multimodal_params] * batch_size # This hidden-commit prefill path does not consume token ids, but # InferState uses input_ids.shape[0] to build position ids. Keep it # aligned with the fixed-shape target verify batch. @@ -217,5 +217,6 @@ def build_block_draft_input( draft_input.mem_indexes = draft_mem_indexes_cpu.cuda(non_blocking=True) draft_input.b_mark_shared_group = torch.zeros_like(draft_input.b_req_idx) draft_input.b_mark_shared_group[block_size - 1 :: block_size] = block_size - draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(draft_input.batch_size)] + empty_multimodal_params = {"images": [], "audios": []} + draft_input.multimodal_params = [empty_multimodal_params] * draft_input.batch_size return draft_input, draft_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/speculative/proposers/dspark.py index ca68069835..a282ee2ca8 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dspark.py @@ -55,7 +55,6 @@ def propose_next( return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, - draft_probs=[] if self.enable_dynamic_spec else None, schedule_probs=schedule_probs, ) @@ -87,7 +86,6 @@ def propose_next( block_token_ids = flat_token_ids.reshape(num_reqs, block_size) proposal_token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] - draft_probs = None if self.enable_dynamic_spec: confidence_logits = draft_model_output.confidence_logits if confidence_logits is None: @@ -99,35 +97,16 @@ def propose_next( assert ( confidence_logits.shape[1] == block_size ), f"confidence logits must have {block_size} columns, got {confidence_logits.shape[1]}" - schedule_probs = self._scatter_step_probs( + # Match the clamp used by the GPU dynamic row selector before it + # converts conditional confidence to prefix survival probability. + schedule_probs = self.scatter_selected_step_probs( selected_rows=selected_rows, - probs=confidence_logits[:, :draft_step].sigmoid(), + selected_probs=confidence_logits[:, :draft_step].sigmoid().clamp(min=0.01, max=0.99), verify_row_count=verify_row_count, ) return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=draft_mem_indexes_cpu, - draft_probs=draft_probs, schedule_probs=schedule_probs, ) - - def _scatter_step_probs( - self, - selected_rows: torch.Tensor, - probs: torch.Tensor, - verify_row_count: int, - ): - assert selected_rows.ndim == 1, "selected_rows must be 1D" - assert probs.ndim == 2, "confidence probabilities must be [selected_rows, draft_step]" - assert probs.shape[0] == selected_rows.shape[0], ( - "confidence probability rows must match selected rows: " f"{selected_rows.shape[0]}, got {probs.shape[0]}" - ) - # Keep the async CPU capacity estimate aligned with the GPU dynamic - # selector, which clamps conditional draft probabilities before - # converting them to prefix survival scores. Unselected rows remain - # zero because scatter_selected_step_probs initializes the output. - probs = probs.clamp(min=0.01, max=0.99) - out = probs.new_zeros((verify_row_count, probs.shape[1]), dtype=torch.float32) - out[selected_rows, :] = probs.float() - return out diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py index c0a140e8ad..4e395d2569 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -5,11 +5,10 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer +from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal -class RecurrentEagleMTPProposer(VanillaMTPProposer): +class RecurrentEagleMTPProposer(BaseSpecProposer): """Shared draft-state setup for recurrent Eagle MTP proposers.""" def build_initial_draft_state( @@ -46,9 +45,6 @@ def build_initial_draft_state_overlap( ) self.backend.draft_models[0].microbatch_overlap_prefill(draft_model_input0, draft_model_input1) - def project_draft_decode_hidden(self, draft_hidden: torch.Tensor) -> torch.Tensor: - return draft_hidden - def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return draft_token_ids @@ -290,7 +286,7 @@ def propose_next_overlap( self.make_single_step_decode_input( base_input=model_input, input_ids=draft_next_token_ids[index], - draft_hidden=self.project_draft_decode_hidden(draft_hiddens[index]), + draft_hidden=draft_hiddens[index], b_req_idx=selected_req_idxs[index], b_mtp_index=selected_mtp_idxs[index], b_seq_len=selected_seq_lens[index], @@ -337,24 +333,6 @@ def propose_next( ) -> SpecProposal: assert 0 <= draft_step <= self.backend.max_draft_step assert accept_len is not None - return self._propose_recurrent( - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=draft_step, - accept_len=accept_len, - ) - - def _propose_recurrent( - self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor, - ) -> SpecProposal: assert main_model_output is not None and main_model_output.spec_hidden is not None verify_row_count = int(next_token_ids.shape[0]) num_reqs = int(b_req_mtp_start_loc.shape[0]) @@ -425,7 +403,7 @@ def _propose_recurrent( draft_input = self.make_single_step_decode_input( base_input=main_model_input, input_ids=draft_next_token_ids, - draft_hidden=self.project_draft_decode_hidden(draft_hidden), + draft_hidden=draft_hidden, b_req_idx=selected_req_idx, b_mtp_index=selected_mtp_index, b_seq_len=selected_seq_len, diff --git a/lightllm/server/router/model_infer/speculative/runner.py b/lightllm/server/router/model_infer/speculative/runner.py index a34e9d9b35..7ab72b1fc5 100644 --- a/lightllm/server/router/model_infer/speculative/runner.py +++ b/lightllm/server/router/model_infer/speculative/runner.py @@ -67,12 +67,11 @@ def run_speculative_forward( ) -> SpecDecodeForwardState: engine = self.engine b_req_mtp_start_loc = gen_b_req_mtp_start_loc(model_input.b_mtp_index, num_reqs=req_num) - verify_result = engine.verify_target_tokens( + accept_len, accepted_index = engine.verify_target_tokens( new_next_token_ids=next_token_ids, b_req_idx=model_input.b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, ) - accepted_index = verify_result.accepted_index if self.backend.is_linear_att_mixed_model: linear_att_spec_state_index_update( req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, @@ -96,7 +95,7 @@ def run_speculative_forward( next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=plan.draft_step, - accept_len=verify_result.accept_len, + accept_len=accept_len, ) all_next_token_ids = engine.pad_all_next_token_ids( @@ -120,7 +119,7 @@ def run_speculative_forward( b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=all_next_token_ids, b_req_idx=model_input.b_req_idx, - spec_accept_len=verify_result.accept_len, + spec_accept_len=accept_len, all_next_token_probs=all_next_token_probs, ) @@ -138,7 +137,7 @@ def run_speculative_forward( spec_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="spec_accept_len", - gpu_tensor=verify_result.accept_len, + gpu_tensor=accept_len, ) sync_event = torch.cuda.Event() @@ -172,14 +171,13 @@ def resolve_pre_post_reqs(self, state: SpecDecodeForwardState, decode_reqs: List accepted_index_cpu=state.accepted_index_cpu, ) - def finish_post(self, state: SpecDecodeForwardState, req_num: int, run_reqs: List) -> SpecDecodePostState: + def finish_post(self, state: SpecDecodeForwardState, req_num: int) -> SpecDecodePostState: state.sync_event.synchronize() engine = self.engine if engine.enable_dynamic_spec: engine.update_dynamic_accept_stats( req_num=req_num, - run_reqs=run_reqs, accepted_index_cpu=state.accepted_index_cpu, spec_accept_len_cpu=state.spec_accept_len_cpu, dynamic_batch_size=state.plan.dynamic_batch_size, diff --git a/lightllm/server/router/model_infer/speculative/verifier.py b/lightllm/server/router/model_infer/speculative/verifier.py deleted file mode 100644 index 0272ee346b..0000000000 --- a/lightllm/server/router/model_infer/speculative/verifier.py +++ /dev/null @@ -1,100 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass -from typing import Optional - -import torch - -from lightllm.common.basemodel.triton_kernel.mtp_utils import mtp_scatter_next_token_ids, mtp_verify - - -@dataclass -class SpecVerifyResult: - """GPU verification result produced after target decode. - - `accept_len` is per logical request and includes the target token at - position 0: - - accept_len: [logical_req_num] - - `accepted_index` is per verified row in the target batch. It marks rows - whose token should be committed and post-processed: - - accepted_index: [verify_batch] - """ - - accept_len: torch.Tensor - accepted_index: torch.Tensor - - -class SpecVerifier: - """Service verifier for LightLLM's speculative row layout. - - DeepSpec verifies by running target logits over - [current_token + draft_tokens] and applying rejection sampling. LightLLM's - service path stores pending candidates in req_sampling_params_manager and - uses Triton kernels to: - - compare the target sampled tokens with previously scattered candidates - - compute per-request accepted prefix length - - scatter the newly proposed candidates for the next decode iteration - """ - - def __init__(self, backend) -> None: - self.backend = backend - - def verify_target_tokens( - self, - new_next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - ) -> SpecVerifyResult: - """Verify target sampled ids against previously proposed ids. - - Inputs: - - `new_next_token_ids`: target sampled ids, shape [verify_batch] - - `b_req_idx`: request ids for each target row, shape [verify_batch] - - `b_req_mtp_start_loc`: first row of each logical request in the - speculative verify batch, shape [logical_req_num] - """ - - spec_accept_len, accepted_index = mtp_verify( - req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - new_next_token_ids=new_next_token_ids, - b_req_idx=b_req_idx, - ) - return SpecVerifyResult(accept_len=spec_accept_len, accepted_index=accepted_index) - - def scatter_next_tokens( - self, - b_req_mtp_start_loc: torch.Tensor, - all_next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - spec_accept_len: torch.Tensor, - all_next_token_probs: Optional[torch.Tensor] = None, - ) -> None: - """Scatter target+draft candidates into per-request next-token buffers. - - Inputs: - - `all_next_token_ids`: [verify_batch, max_draft_step + 1]. Column 0 is the - target sampled token from this iteration; remaining columns are draft - candidates padded to the configured speculative width when dynamic scheduling - produces a shorter proposal. - - `all_next_token_probs`: optional [verify_batch, max_draft_step + 1]. - Dynamic scheduling stores selected-token probabilities, not full - vocab distributions. - """ - - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=b_req_idx, - spec_accept_len=spec_accept_len, - req_to_next_token_probs=( - self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs - if all_next_token_probs is not None - else None - ), - all_next_token_probs=all_next_token_probs, - ) diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index bca619dfe9..429e7c60b3 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -6,6 +6,7 @@ from lightllm.common.basemodel import batch_objs from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel from lightllm.models.qwen3_dspark import model_output as dspark_model_output +from lightllm.models.qwen3_dspark.layer_weights import pre_and_post_layer_weight as dspark_pre_post_weight from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput @@ -60,6 +61,39 @@ def test_dspark_no_ref_conversion_dispatches_to_dspark_fields(monkeypatch): assert all(converted != original for converted, original in zip(converted_ptrs, original_ptrs)) +def test_confidence_head_does_not_inherit_model_quantization(monkeypatch): + captured_kwargs = {} + + def init_base_weight(self, data_type, network_config, quant_cfg): + self.data_type_ = data_type + self.quant_cfg = quant_cfg + + class RecordingROWMMWeight: + def __init__(self, **kwargs): + captured_kwargs.update(kwargs) + + monkeypatch.setattr( + dspark_pre_post_weight.Qwen3DFlashPreAndPostLayerWeight, + "__init__", + init_base_weight, + ) + monkeypatch.setattr(dspark_pre_post_weight, "ROWMMWeight", RecordingROWMMWeight) + + quant_method = object() + quant_cfg = SimpleNamespace(get_quant_method=lambda *_: quant_method) + dspark_pre_post_weight.Qwen3DSparkPreAndPostLayerWeight( + data_type=torch.bfloat16, + network_config={ + "hidden_size": 16, + "vocab_size": 32, + "enable_confidence_head": True, + }, + quant_cfg=quant_cfg, + ) + + assert captured_kwargs["quant_method"] is None + + def test_vanilla_markov_local_sampling_matches_full_logits(): post_infer = Qwen3DSparkPostLayerInfer.__new__(Qwen3DSparkPostLayerInfer) post_infer.block_size_ = 3 From f5293ba826eef53eca4b2789cc6697417cd4f264 Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Tue, 11 Aug 2026 11:21:24 +0000 Subject: [PATCH 008/103] fix: align speculative draft inputs and DSpark RoPE --- lightllm/models/qwen3_5_dspark/model.py | 9 ++- .../mode_backend/chunked_prefill/impl.py | 4 +- .../mode_backend/dp_backend/impl.py | 35 +++------ .../router/model_infer/speculative/engine.py | 65 +++-------------- .../model_infer/speculative/proposers/base.py | 43 +++-------- .../speculative/proposers/dflash.py | 12 +-- .../speculative/proposers/eagle_mtp.py | 43 ++++++----- .../speculative/proposers/vanilla_mtp.py | 73 +++++++++---------- .../models/test_qwen3_dspark_model_output.py | 24 ++++++ 9 files changed, 130 insertions(+), 178 deletions(-) diff --git a/lightllm/models/qwen3_5_dspark/model.py b/lightllm/models/qwen3_5_dspark/model.py index 68cd359340..cec8b7b8c9 100644 --- a/lightllm/models/qwen3_5_dspark/model.py +++ b/lightllm/models/qwen3_5_dspark/model.py @@ -14,10 +14,11 @@ def _init_config(self): rope_parameters = self.config["rope_parameters"] if "rope_theta" in rope_parameters and "rope_theta" not in self.config: self.config["rope_theta"] = rope_parameters["rope_theta"] - if "partial_rotary_factor" in rope_parameters and "partial_rotary_factor" not in self.config: - self.config["partial_rotary_factor"] = rope_parameters["partial_rotary_factor"] - if "rope_scaling" not in self.config: - self.config["rope_scaling"] = rope_parameters + + # DeepSpec trains this Qwen3 draft with ordinary 1D full-head RoPE. Do + # not pass the Qwen3.5 target's MRoPE layout into the draft backbone. + self.config["rope_scaling"] = None + self.config["partial_rotary_factor"] = 1.0 def _init_custom(self): # Draft and target use different rotary shapes, so the draft owns its rotary cache. diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 8e42f73013..50479ee0e3 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -207,8 +207,8 @@ def prefill_mtp( # mtp kv fill spec_engine = self.spec_engine spec_engine.build_initial_draft_state( - model_input=model_input, - model_output=model_output, + target_model_input=model_input, + target_model_output=model_output, next_token_ids=next_token_ids, ) g_infer_context.copy_linear_att_state_to_cache_buffer( diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 458bf07e0b..e04f70eba3 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -463,8 +463,8 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] copy_len=req_num, ) self.spec_engine.build_initial_draft_state( - model_input=model_input, - model_output=model_output, + target_model_input=model_input, + target_model_output=model_output, next_token_ids=draft_next_token_ids_gpu, ) if req_num > 0: @@ -636,12 +636,8 @@ def _draft_decode_vanilla( # process the draft model output for draft_model_idx in range(self.max_draft_step): - - draft_model_input = self.spec_engine.prepare_draft_decode_input( - model_input=draft_model_input, - next_token_ids=draft_next_token_ids_gpu, - mtp_draft_input_hiddens=draft_hidden, - ) + draft_model_input.input_ids = draft_next_token_ids_gpu + draft_model_input.mtp_draft_input_hiddens = draft_hidden # spec decode: MTP draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) draft_hidden = draft_model_output.spec_hidden @@ -772,11 +768,11 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I ) self.spec_engine.build_initial_draft_state_overlap( - model_input0=model_input0, - model_output0=model_output0, + target_model_input0=model_input0, + target_model_output0=model_output0, next_token_ids0=draft_next_token_ids_gpu0, - model_input1=model_input1, - model_output1=model_output1, + target_model_input1=model_input1, + target_model_output1=model_output1, next_token_ids1=draft_next_token_ids_gpu1, ) @@ -978,17 +974,10 @@ def _draft_decode_vanilla_overlap( # process the draft model output for draft_model_idx in range(self.max_draft_step): - - draft_model_input0 = self.spec_engine.prepare_draft_decode_input( - model_input=draft_model_input0, - next_token_ids=draft_next_token_ids_gpu0, - mtp_draft_input_hiddens=draft_hidden0, - ) - draft_model_input1 = self.spec_engine.prepare_draft_decode_input( - model_input=draft_model_input1, - next_token_ids=draft_next_token_ids_gpu1, - mtp_draft_input_hiddens=draft_hidden1, - ) + draft_model_input0.input_ids = draft_next_token_ids_gpu0 + draft_model_input0.mtp_draft_input_hiddens = draft_hidden0 + draft_model_input1.input_ids = draft_next_token_ids_gpu1 + draft_model_input1.mtp_draft_input_hiddens = draft_hidden1 draft_model_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( draft_model_input0, draft_model_input1 diff --git a/lightllm/server/router/model_infer/speculative/engine.py b/lightllm/server/router/model_infer/speculative/engine.py index 2f8cc05cbc..0ebe83bd8e 100644 --- a/lightllm/server/router/model_infer/speculative/engine.py +++ b/lightllm/server/router/model_infer/speculative/engine.py @@ -63,74 +63,33 @@ def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) return g_infer_context.req_manager.mem_manager.alloc(token_count) - def prepare_draft_prefill_input( - self, - model_input: ModelInput, - next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: torch.Tensor, - ) -> ModelInput: - """Build draft prefill input from target prefill input. - - `next_token_ids`: [run_req_num] - `mtp_draft_input_hiddens`: captured target feature. The first - dimension matches the target prefill token layout after padding/unpad - handling; the second dimension is either hidden_size or - hidden_size * len(target_layer_ids). - """ - - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - return prepare_mtp_prefill_inputs( - model_input=model_input, - b_next_token_ids=next_token_ids, - mtp_draft_input_hiddens=mtp_draft_input_hiddens, - ) - - def prepare_draft_decode_input( - self, - model_input: ModelInput, - next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: torch.Tensor, - ) -> ModelInput: - """Mutate a decode ModelInput for one draft forward. - - `next_token_ids`: [verify_batch] - `mtp_draft_input_hiddens`: [verify_batch, hidden_dim_for_draft] - """ - - model_input.input_ids = next_token_ids - model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens - if self.spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3"): - model_input.draft_step = 0 - return model_input - def build_initial_draft_state( self, - model_input: ModelInput, - model_output: ModelOutput, + target_model_input: ModelInput, + target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: self.proposer.build_initial_draft_state( - model_input=model_input, - model_output=model_output, + target_model_input=target_model_input, + target_model_output=target_model_output, next_token_ids=next_token_ids, ) def build_initial_draft_state_overlap( self, - model_input0: ModelInput, - model_output0: ModelOutput, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, next_token_ids0: torch.Tensor, - model_input1: ModelInput, - model_output1: ModelOutput, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: self.proposer.build_initial_draft_state_overlap( - model_input0=model_input0, - model_output0=model_output0, + target_model_input0=target_model_input0, + target_model_output0=target_model_output0, next_token_ids0=next_token_ids0, - model_input1=model_input1, - model_output1=model_output1, + target_model_input1=target_model_input1, + target_model_output1=target_model_output1, next_token_ids1=next_token_ids1, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/base.py b/lightllm/server/router/model_infer/speculative/proposers/base.py index 2a0b565556..c4d1ff1d80 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/speculative/proposers/base.py @@ -68,30 +68,6 @@ def __init__(self, engine: "SpecEngine") -> None: def enable_dynamic_spec(self) -> bool: return self.engine.enable_dynamic_spec - def prepare_draft_prefill_input( - self, - model_input: ModelInput, - next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: torch.Tensor, - ) -> ModelInput: - return self.engine.prepare_draft_prefill_input( - model_input=model_input, - next_token_ids=next_token_ids, - mtp_draft_input_hiddens=mtp_draft_input_hiddens, - ) - - def prepare_draft_decode_input( - self, - model_input: ModelInput, - next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: torch.Tensor, - ) -> ModelInput: - return self.engine.prepare_draft_decode_input( - model_input=model_input, - next_token_ids=next_token_ids, - mtp_draft_input_hiddens=mtp_draft_input_hiddens, - ) - def select_accepted_tail_rows(self, b_req_mtp_start_loc: torch.Tensor, accept_len: torch.Tensor) -> torch.Tensor: return (b_req_mtp_start_loc + accept_len - 1).to(torch.long) @@ -116,20 +92,19 @@ def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: def build_initial_draft_state( self, - model_input: ModelInput, - model_output: ModelOutput, + target_model_input: ModelInput, + target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: """Build initial draft KV/state before the first decode verify step. Inputs: - - `model_input`: target prompt ModelInput. Its request order and + - `target_model_input`: target prompt ModelInput. Its request order and mem_indexes are reused by the draft state builder. + - `target_model_output`: target output containing the features needed + by the selected speculative algorithm. - `next_token_ids`: first accepted target token, shape [run_req_num]. - SpecEngine.prepare_draft_prefill_input injects captured target hidden - features into `mtp_draft_input_hiddens`. - This hook only prepares draft-side state. It intentionally does not scatter proposal tokens; the first decode iteration verifies as having no draft candidates and produces the first proposal through @@ -140,11 +115,11 @@ def build_initial_draft_state( def build_initial_draft_state_overlap( self, - model_input0: ModelInput, - model_output0: ModelOutput, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, next_token_ids0: torch.Tensor, - model_input1: ModelInput, - model_output1: ModelOutput, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: """Build initial draft state for two overlapped prefill microbatches.""" diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index e03035021d..fd5128ffc9 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -25,19 +25,19 @@ class DFlashProposer(BaseSpecProposer): @torch.no_grad() def build_initial_draft_state( self, - model_input: ModelInput, - model_output: ModelOutput, + target_model_input: ModelInput, + target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - target_hidden = model_output.spec_hidden + target_hidden = target_model_output.spec_hidden assert target_hidden is not None if target_hidden.numel() == 0: return draft_model = self.backend.draft_models[0] - assert model_input.input_ids is not None - assert model_input.input_ids.shape[0] == target_hidden.shape[0] - draft_input = copy.copy(model_input) + assert target_model_input.input_ids is not None + assert target_model_input.input_ids.shape[0] == target_hidden.shape[0] + draft_input = copy.copy(target_model_input) # DFlash consumes target hidden states directly on this prefill path. draft_input.mtp_draft_input_hiddens = target_hidden draft_model.forward(draft_input) diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py index 4e395d2569..4e71fb1679 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -13,35 +13,42 @@ class RecurrentEagleMTPProposer(BaseSpecProposer): def build_initial_draft_state( self, - model_input: ModelInput, - model_output: ModelOutput, + target_model_input: ModelInput, + target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - draft_model_input = self.prepare_draft_prefill_input( - model_input=model_input, - next_token_ids=next_token_ids, - mtp_draft_input_hiddens=model_output.spec_hidden, + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + assert target_model_output.spec_hidden is not None + draft_model_input = prepare_mtp_prefill_inputs( + model_input=target_model_input, + b_next_token_ids=next_token_ids, + mtp_draft_input_hiddens=target_model_output.spec_hidden, ) self.backend.draft_models[0].forward(draft_model_input) def build_initial_draft_state_overlap( self, - model_input0: ModelInput, - model_output0: ModelOutput, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, next_token_ids0: torch.Tensor, - model_input1: ModelInput, - model_output1: ModelOutput, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - draft_model_input0 = self.prepare_draft_prefill_input( - model_input=model_input0, - next_token_ids=next_token_ids0, - mtp_draft_input_hiddens=model_output0.spec_hidden, + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + assert target_model_output0.spec_hidden is not None + assert target_model_output1.spec_hidden is not None + draft_model_input0 = prepare_mtp_prefill_inputs( + model_input=target_model_input0, + b_next_token_ids=next_token_ids0, + mtp_draft_input_hiddens=target_model_output0.spec_hidden, ) - draft_model_input1 = self.prepare_draft_prefill_input( - model_input=model_input1, - next_token_ids=next_token_ids1, - mtp_draft_input_hiddens=model_output1.spec_hidden, + draft_model_input1 = prepare_mtp_prefill_inputs( + model_input=target_model_input1, + b_next_token_ids=next_token_ids1, + mtp_draft_input_hiddens=target_model_output1.spec_hidden, ) self.backend.draft_models[0].microbatch_overlap_prefill(draft_model_input0, draft_model_input1) diff --git a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py index 7654960259..db0fd974be 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py @@ -21,62 +21,62 @@ class VanillaMTPProposer(BaseSpecProposer): def build_initial_draft_state( self, - model_input: ModelInput, - model_output: ModelOutput, + target_model_input: ModelInput, + target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - draft_model_input = model_input + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + draft_model_input = target_model_input + source_model_output = target_model_output draft_next_token_ids = next_token_ids - draft_hidden = model_output.spec_hidden - assert draft_hidden is not None for draft_model in self.backend.draft_models: - draft_model_input = self.prepare_draft_prefill_input( + assert source_model_output.spec_hidden is not None + draft_model_input = prepare_mtp_prefill_inputs( model_input=draft_model_input, - next_token_ids=draft_next_token_ids, - mtp_draft_input_hiddens=draft_hidden, + b_next_token_ids=draft_next_token_ids, + mtp_draft_input_hiddens=source_model_output.spec_hidden, ) - draft_model_output = draft_model.forward(draft_model_input) - draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) - draft_hidden = draft_model_output.spec_hidden - assert draft_hidden is not None + source_model_output = draft_model.forward(draft_model_input) + draft_next_token_ids = self.backend._gen_argmax_token_ids(source_model_output) def build_initial_draft_state_overlap( self, - model_input0: ModelInput, - model_output0: ModelOutput, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, next_token_ids0: torch.Tensor, - model_input1: ModelInput, - model_output1: ModelOutput, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - draft_model_input0 = model_input0 - draft_model_input1 = model_input1 + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + draft_model_input0 = target_model_input0 + draft_model_input1 = target_model_input1 + source_model_output0 = target_model_output0 + source_model_output1 = target_model_output1 draft_next_token_ids0 = next_token_ids0 draft_next_token_ids1 = next_token_ids1 - draft_hidden0 = model_output0.spec_hidden - draft_hidden1 = model_output1.spec_hidden - assert draft_hidden0 is not None and draft_hidden1 is not None for draft_model in self.backend.draft_models: - draft_model_input0 = self.prepare_draft_prefill_input( + assert source_model_output0.spec_hidden is not None + assert source_model_output1.spec_hidden is not None + draft_model_input0 = prepare_mtp_prefill_inputs( model_input=draft_model_input0, - next_token_ids=draft_next_token_ids0, - mtp_draft_input_hiddens=draft_hidden0, + b_next_token_ids=draft_next_token_ids0, + mtp_draft_input_hiddens=source_model_output0.spec_hidden, ) - draft_model_input1 = self.prepare_draft_prefill_input( + draft_model_input1 = prepare_mtp_prefill_inputs( model_input=draft_model_input1, - next_token_ids=draft_next_token_ids1, - mtp_draft_input_hiddens=draft_hidden1, + b_next_token_ids=draft_next_token_ids1, + mtp_draft_input_hiddens=source_model_output1.spec_hidden, ) - draft_model_output0, draft_model_output1 = draft_model.microbatch_overlap_prefill( + source_model_output0, source_model_output1 = draft_model.microbatch_overlap_prefill( draft_model_input0, draft_model_input1, ) - draft_next_token_ids0 = self.backend._gen_argmax_token_ids(draft_model_output0) - draft_next_token_ids1 = self.backend._gen_argmax_token_ids(draft_model_output1) - draft_hidden0 = draft_model_output0.spec_hidden - draft_hidden1 = draft_model_output1.spec_hidden - assert draft_hidden0 is not None and draft_hidden1 is not None + draft_next_token_ids0 = self.backend._gen_argmax_token_ids(source_model_output0) + draft_next_token_ids1 = self.backend._gen_argmax_token_ids(source_model_output1) def propose_next( self, @@ -96,11 +96,8 @@ def propose_next( for step in range(draft_step): draft_model = self.backend.draft_models[step] - draft_model_input = self.prepare_draft_decode_input( - model_input=draft_model_input, - next_token_ids=draft_next_token_ids, - mtp_draft_input_hiddens=draft_hidden, - ) + draft_model_input.input_ids = draft_next_token_ids + draft_model_input.mtp_draft_input_hiddens = draft_hidden draft_model_output = draft_model.forward(draft_model_input) draft_hidden = draft_model_output.spec_hidden assert draft_hidden is not None diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index 429e7c60b3..b01108c95c 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -94,6 +94,30 @@ def __init__(self, **kwargs): assert captured_kwargs["quant_method"] is None +def test_qwen35_dspark_uses_training_rope_layout(monkeypatch): + def init_dspark_config(self): + self.config = { + "dflash_config": {"mask_token_id": 1}, + "rope_parameters": { + "rope_theta": 1_000_000, + "partial_rotary_factor": 0.25, + "mrope_interleaved": True, + "mrope_section": [11, 11, 10], + "rope_type": "default", + }, + } + + monkeypatch.setattr(Qwen3DSparkModel, "_init_config", init_dspark_config) + model = Qwen3_5DSparkModel.__new__(Qwen3_5DSparkModel) + + model._init_config() + + assert model.config["rope_scaling"] is None + assert model.config["rope_theta"] == 1_000_000 + assert model.config["partial_rotary_factor"] == 1.0 + assert model.config["mask_token_id"] == 1 + + def test_vanilla_markov_local_sampling_matches_full_logits(): post_infer = Qwen3DSparkPostLayerInfer.__new__(Qwen3DSparkPostLayerInfer) post_infer.block_size_ = 3 From f713b82745c074acb29f63959992840870604ce3 Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Thu, 13 Aug 2026 08:51:53 +0000 Subject: [PATCH 009/103] refactor: simplify dynamic speculative scheduling --- .../triton_kernel/dynamic_spec_utils.py | 271 +--- .../basemodel/triton_kernel/mtp_utils.py | 91 +- lightllm/common/req_manager.py | 8 +- lightllm/models/qwen3_dflash/model.py | 2 - .../layer_infer/post_layer_infer.py | 1 - .../pre_and_post_layer_weight.py | 5 +- lightllm/models/qwen3_dspark/model.py | 6 - lightllm/models/qwen3_eagle/model.py | 2 - lightllm/server/api_cli.py | 3 +- lightllm/server/router/manager.py | 2 +- .../model_infer/mode_backend/base_backend.py | 45 +- .../mode_backend/chunked_prefill/impl.py | 124 +- .../mode_backend/dp_backend/impl.py | 119 +- .../mode_backend/generic_post_process.py | 72 +- .../mode_backend/generic_pre_process.py | 7 +- .../router/model_infer/pin_mem_manager.py | 20 + .../router/model_infer/speculative/engine.py | 455 +++---- .../router/model_infer/speculative/planner.py | 1189 ++++------------- .../speculative/proposers/__init__.py | 13 +- .../model_infer/speculative/proposers/base.py | 95 +- .../speculative/proposers/dflash.py | 60 +- .../speculative/proposers/dspark.py | 51 +- .../speculative/proposers/eagle3.py | 78 +- .../speculative/proposers/eagle_mtp.py | 176 ++- .../speculative/proposers/vanilla_mtp.py | 36 +- .../router/model_infer/speculative/runner.py | 207 --- test/test_api/test_gsmk.py | 42 +- .../triton_kernel/test_dynamic_spec_utils.py | 118 +- .../basemodel/triton_kernel/test_mtp_utils.py | 46 +- .../models/test_qwen3_dspark_model_output.py | 117 +- .../mode_backend/test_dp_spec_engine.py | 13 +- .../mode_backend/test_generic_post_process.py | 50 - .../speculative/test_eagle_overlap.py | 25 +- .../model_infer/speculative/test_planner.py | 460 ++++++- unit_tests/utils/test_speculative_utils.py | 15 +- 35 files changed, 1684 insertions(+), 2340 deletions(-) delete mode 100644 lightllm/server/router/model_infer/speculative/runner.py delete mode 100644 unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py index a54aad84ab..74bc9462f1 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py @@ -1,139 +1,43 @@ import triton import triton.language as tl -from triton.language.standard import _log2, sum, zeros_like import torch @triton.jit -def _fwd_kernel_cumprod_probs( - req_to_next_token_probs, - req_to_next_token_probs_stride, +def _fwd_kernel_cumprod_scores( + req_to_next_token_scores, + req_to_next_token_scores_stride, b_req_idx, max_draft_step, BLOCK_SIZE: tl.constexpr, ): cur_index = tl.program_id(0) cur_req_idx = tl.load(b_req_idx + cur_index * (max_draft_step + 1)) - base_ptr = req_to_next_token_probs + cur_req_idx * req_to_next_token_probs_stride + base_ptr = req_to_next_token_scores + cur_req_idx * req_to_next_token_scores_stride tl.store(base_ptr, 1.0) offset = tl.arange(0, BLOCK_SIZE) store_mask = offset < (max_draft_step + 1) - probs = tl.load(base_ptr + offset, mask=store_mask, other=0.0) + scores = tl.load(base_ptr + offset, mask=store_mask, other=0.0) # offset 0 是 target sample,本轮恒接受;只有 draft 条件接受概率需要 clamp。 - probs = tl.where(offset == 0, 1.0, probs) - # 对于 draft probs 中大于 0.99 的值,设置为 0.99,避免错误的值,照成后续的采样操作失败。 - # 对于 draft probs 中小于 0.01 的值,设置为 0.01,避免错误的值,照成后续的采样操作失败。 + scores = tl.where(offset == 0, 1.0, scores) + # Clamp draft scheduling scores before converting them to prefix-survival scores. # This makes each request's cumulative acceptance probabilities monotonic, # so global top-k selection cannot pick a later draft row without its prefix. - probs = tl.where((offset != 0) & (probs >= 0.99), 0.99, probs) - probs = tl.where((offset != 0) & (probs <= 0.01), 0.01, probs) + scores = tl.where((offset != 0) & (scores >= 0.99), 0.99, scores) + scores = tl.where((offset != 0) & (scores <= 0.01), 0.01, scores) - cum_probs = tl.cumprod(probs, axis=0) + cumulative_scores = tl.cumprod(scores, axis=0) - tl.store(base_ptr + offset, cum_probs, mask=store_mask) - return - - -@triton.jit -def _compare_and_swap(x, ids, flip, i: tl.core.constexpr, n_dims: tl.core.constexpr): - n_outer: tl.core.constexpr = x.numel >> n_dims - shape: tl.core.constexpr = [n_outer * 2 ** i, 2, 2 ** (n_dims - i - 1)] - y = tl.core.reshape(x, shape) - # slice left/right with 'stride' 2**(n_dims - i - 1) - mask = tl.core.arange(0, 2)[None, :, None] - left = tl.core.broadcast_to(sum(y * (1 - mask), 1)[:, None, :], shape) - right = tl.core.broadcast_to(sum(y * mask, 1)[:, None, :], shape) - left = tl.core.reshape(left, x.shape) - right = tl.core.reshape(right, x.shape) - - y_idx = tl.core.reshape(ids, shape) - left_idx = tl.core.broadcast_to(sum(y_idx * (1 - mask), 1)[:, None, :], shape) - right_idx = tl.core.broadcast_to(sum(y_idx * mask, 1)[:, None, :], shape) - left_idx = tl.core.reshape(left_idx, x.shape) - right_idx = tl.core.reshape(right_idx, x.shape) - - idtype = tl.core.get_int_dtype(bitwidth=x.dtype.primitive_bitwidth, signed=True) - ileft = left.to(idtype, bitcast=True) - iright = right.to(idtype, bitcast=True) - ix = x.to(idtype, bitcast=True) - - cond = (left > right) != (flip != 0) - - ret = ix ^ tl.core.where(cond, ileft ^ iright, zeros_like(ix)) - new_ids = ids ^ tl.core.where(cond, left_idx ^ right_idx, zeros_like(ids)) - - return ret.to(x.dtype, bitcast=True), new_ids - - -@triton.jit -def _bitonic_merge(x, ids, stage: tl.core.constexpr, order: tl.core.constexpr, n_dims: tl.core.constexpr): - """ - order_type 0 == ascending - order_type 1 == descending - order_type 2 == alternating - """ - n_outer: tl.core.constexpr = x.numel >> n_dims - tl.core.static_assert(stage <= n_dims) - if order == 2: - shape: tl.core.constexpr = [n_outer * 2 ** (n_dims - 1 - stage), 2, 2 ** stage] - flip = tl.core.reshape(tl.core.broadcast_to(tl.core.arange(0, 2)[None, :, None], shape), x.shape) - else: - flip = order - for i in tl.core.static_range(stage): - x, ids = _compare_and_swap(x, ids, flip, i + (n_dims - stage), n_dims) - return x, ids - - -@triton.jit -def argsort(x, ids, dim: tl.core.constexpr = None, descending: tl.core.constexpr = tl.core.CONSTEXPR_0): - _dim: tl.core.constexpr = len(x.shape) - 1 if dim is None else dim - tl.core.static_assert(_dim == len(x.shape) - 1, "only minor dimension is currently supported") - n_dims: tl.core.constexpr = _log2(x.shape[_dim]) - - for i in tl.core.static_range(1, n_dims + 1): - x, ids = _bitonic_merge(x, ids, i, 2 if i < n_dims else descending, n_dims) - return x, ids - - -@triton.jit -def _fwd_kernel_select_dynamic_spec_rows( - req_to_next_token_probs, - req_to_next_token_probs_stride, - selected_row_mask, - b_req_idx, - max_draft_step, - pre_draft_step, - req_num, - dynamic_batch_size, - BLOCK_SIZE: tl.constexpr, -): - all_num = req_num * (pre_draft_step + 1) - offset = tl.arange(0, BLOCK_SIZE) - mask = offset < all_num - req_offset = offset // (pre_draft_step + 1) - next_token_offset = offset % (pre_draft_step + 1) - original_offset = req_offset * (max_draft_step + 1) + next_token_offset - - req_idx_index = tl.load(b_req_idx + original_offset, mask=mask, other=0) - - probs = tl.load( - req_to_next_token_probs + req_idx_index * req_to_next_token_probs_stride + next_token_offset, - mask=mask, - other=-1.0, - ) - - sorted_probs, sorted_ids = argsort(probs, original_offset, descending=True) - - tl.store(selected_row_mask + sorted_ids, 1, mask=mask & (offset < dynamic_batch_size)) + tl.store(base_ptr + offset, cumulative_scores, mask=store_mask) return def sample_dynamic_spec_row_mask( dynamic_batch_size: int, b_req_idx: torch.Tensor, - req_to_next_token_probs: torch.Tensor, + req_to_next_token_scores: torch.Tensor, max_draft_step: int, pre_draft_step: int = None, ) -> torch.Tensor: @@ -142,16 +46,16 @@ def sample_dynamic_spec_row_mask( pre_draft_step = max_draft_step if pre_draft_step is None else int(pre_draft_step) assert 0 <= pre_draft_step <= max_draft_step assert b_req_idx.shape[0] % (max_draft_step + 1) == 0 - assert req_to_next_token_probs.is_cuda + assert req_to_next_token_scores.is_cuda assert dynamic_batch_size <= b_req_idx.shape[0] req_num = len(b_req_idx) // (max_draft_step + 1) valid_row_num = req_num * (pre_draft_step + 1) assert dynamic_batch_size <= valid_row_num - # cumprod probs for each request - _fwd_kernel_cumprod_probs[(req_num,)]( - req_to_next_token_probs=req_to_next_token_probs, - req_to_next_token_probs_stride=req_to_next_token_probs.stride(0), + # Convert each request's conditional scheduling scores to prefix survival scores. + _fwd_kernel_cumprod_scores[(req_num,)]( + req_to_next_token_scores=req_to_next_token_scores, + req_to_next_token_scores_stride=req_to_next_token_scores.stride(0), b_req_idx=b_req_idx, max_draft_step=max_draft_step, BLOCK_SIZE=triton.next_power_of_2(max_draft_step + 1), @@ -159,138 +63,13 @@ def sample_dynamic_spec_row_mask( num_stages=1, ) - # 1 为选中, 0 为未选中 - selected_row_mask = torch.zeros((len(b_req_idx),), dtype=torch.int32, device="cuda") + request_ids = b_req_idx[:: max_draft_step + 1].long() + scores = req_to_next_token_scores.index_select(0, request_ids)[:, : pre_draft_step + 1].flatten() + compact_ids = torch.topk(scores, k=dynamic_batch_size, sorted=False).indices + request_offsets = compact_ids // (pre_draft_step + 1) + step_offsets = compact_ids % (pre_draft_step + 1) + selected_ids = request_offsets * (max_draft_step + 1) + step_offsets - grid = (1,) - _fwd_kernel_select_dynamic_spec_rows[grid]( - req_to_next_token_probs=req_to_next_token_probs, - req_to_next_token_probs_stride=req_to_next_token_probs.stride(0), - selected_row_mask=selected_row_mask, - b_req_idx=b_req_idx, - max_draft_step=max_draft_step, - pre_draft_step=pre_draft_step, - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - BLOCK_SIZE=triton.next_power_of_2(valid_row_num), - num_warps=1, - num_stages=1, - ) + selected_row_mask = torch.zeros((len(b_req_idx),), dtype=torch.int32, device=b_req_idx.device) + selected_row_mask.scatter_(0, selected_ids, 1) return selected_row_mask - - -@triton.jit -def _fwd_kernel_trim_post_sample_tensors( - b_req_idx, - out_b_req_idx, - b_temperatures, - out_b_temperatures, - b_top_ps, - out_b_top_ps, - b_top_ks, - out_b_top_ks, - b_length_penalty_param, - out_b_length_penalty_param, - b_mask_eos_reqs, - out_b_mask_eos_reqs, - selected_row_mask, - selected_dst_pos, - batch_size, - BLOCK_SIZE: tl.constexpr, -): - offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offsets < batch_size - selected = tl.load(selected_row_mask + offsets, mask=mask, other=0) != 0 - dst_pos = tl.load(selected_dst_pos + offsets, mask=mask, other=0) - write_mask = mask & selected - - tl.store( - out_b_req_idx + dst_pos, - tl.load(b_req_idx + offsets, mask=mask, other=0), - mask=write_mask, - ) - tl.store( - out_b_temperatures + dst_pos, - tl.load(b_temperatures + offsets, mask=mask, other=0.0), - mask=write_mask, - ) - tl.store( - out_b_top_ps + dst_pos, - tl.load(b_top_ps + offsets, mask=mask, other=0.0), - mask=write_mask, - ) - tl.store( - out_b_top_ks + dst_pos, - tl.load(b_top_ks + offsets, mask=mask, other=0), - mask=write_mask, - ) - tl.store( - out_b_length_penalty_param + dst_pos, - tl.load(b_length_penalty_param + offsets, mask=mask, other=0), - mask=write_mask, - ) - tl.store( - out_b_mask_eos_reqs + dst_pos, - tl.load(b_mask_eos_reqs + offsets, mask=mask, other=0), - mask=write_mask, - ) - - -def trim_post_sample_tensors( - dynamic_batch_size: int, - selected_row_mask: torch.Tensor, - b_req_idx: torch.Tensor, - b_temperatures: torch.Tensor, - b_top_ps: torch.Tensor, - b_top_ks: torch.Tensor, - b_length_penalty_param: torch.Tensor, - b_mask_eos_reqs: torch.Tensor, -): - assert selected_row_mask.is_cuda - dynamic_batch_size = int(dynamic_batch_size) - selected_row_mask = selected_row_mask.to(torch.int32) - selected_dst_pos = torch.cumsum(selected_row_mask, dim=0, dtype=torch.int32) - 1 - batch_size = selected_row_mask.shape[0] - - output_tensors = tuple( - torch.empty((dynamic_batch_size,), dtype=tensor.dtype, device=tensor.device) - for tensor in ( - b_req_idx, - b_temperatures, - b_top_ps, - b_top_ks, - b_length_penalty_param, - b_mask_eos_reqs, - ) - ) - ( - out_b_req_idx, - out_b_temperatures, - out_b_top_ps, - out_b_top_ks, - out_b_length_penalty_param, - out_b_mask_eos_reqs, - ) = output_tensors - - block_size = 256 - _fwd_kernel_trim_post_sample_tensors[(triton.cdiv(batch_size, block_size),)]( - b_req_idx=b_req_idx, - out_b_req_idx=out_b_req_idx, - b_temperatures=b_temperatures, - out_b_temperatures=out_b_temperatures, - b_top_ps=b_top_ps, - out_b_top_ps=out_b_top_ps, - b_top_ks=b_top_ks, - out_b_top_ks=out_b_top_ks, - b_length_penalty_param=b_length_penalty_param, - out_b_length_penalty_param=out_b_length_penalty_param, - b_mask_eos_reqs=b_mask_eos_reqs, - out_b_mask_eos_reqs=out_b_mask_eos_reqs, - selected_row_mask=selected_row_mask, - selected_dst_pos=selected_dst_pos, - batch_size=batch_size, - BLOCK_SIZE=block_size, - num_warps=4, - num_stages=1, - ) - return output_tensors diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index f411db232c..8cee6e5163 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -98,15 +98,16 @@ def _fwd_kernel_mtp_scatter_next_token_ids( req_to_next_token_ids_stride, all_next_token_ids, all_next_token_ids_stride, - req_to_next_token_probs, - req_to_next_token_probs_stride, - all_next_token_probs, - all_next_token_probs_stride, + req_to_next_token_scores, + req_to_next_token_scores_stride, + schedule_scores, + schedule_scores_stride, spec_accept_len, b_req_mtp_start_loc, b_req_idx, - mtp_step, - HAS_NEXT_TOKEN_PROBS: tl.constexpr, + proposal_width, + verify_width, + HAS_NEXT_TOKEN_SCORES: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): @@ -115,27 +116,32 @@ def _fwd_kernel_mtp_scatter_next_token_ids( accept_len = tl.load(spec_accept_len + cur_index) cur_req_idx = tl.load(b_req_idx + req_start_loc) offset = tl.arange(0, BLOCK_SIZE) - - if HAS_NEXT_TOKEN_PROBS: - cur_next_token_probs = tl.load( - all_next_token_probs + (req_start_loc + accept_len - 1) * all_next_token_probs_stride + offset, - mask=offset < mtp_step, + selected_row = req_start_loc + accept_len - 1 + + if HAS_NEXT_TOKEN_SCORES: + # schedule_scores omits the guaranteed target column. Insert its 1.0 + # here and clear the unused tail of the fixed-width request buffer. + schedule_offset = tl.maximum(offset - 1, 0) + draft_scores = tl.load( + schedule_scores + selected_row * schedule_scores_stride + schedule_offset, + mask=(offset > 0) & (offset < proposal_width), other=0.0, ) + next_token_scores = tl.where(offset == 0, 1.0, draft_scores) tl.store( - req_to_next_token_probs + cur_req_idx * req_to_next_token_probs_stride + offset, - cur_next_token_probs, - mask=offset < mtp_step, + req_to_next_token_scores + cur_req_idx * req_to_next_token_scores_stride + offset, + next_token_scores, + mask=offset < verify_width, ) scatter_next_token_ids = tl.load( - all_next_token_ids + (req_start_loc + accept_len - 1) * all_next_token_ids_stride + offset, - mask=offset < mtp_step, - other=0, + all_next_token_ids + selected_row * all_next_token_ids_stride + offset, + mask=offset < proposal_width, + other=1, ) tl.store( req_to_next_token_ids + cur_req_idx * req_to_next_token_ids_stride + offset, scatter_next_token_ids, - mask=offset < mtp_step, + mask=offset < verify_width, ) return @@ -146,30 +152,30 @@ def mtp_scatter_next_token_ids( all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, spec_accept_len: torch.Tensor, - req_to_next_token_probs: Optional[torch.Tensor] = None, - all_next_token_probs: Optional[torch.Tensor] = None, + req_to_next_token_scores: Optional[torch.Tensor] = None, + schedule_scores: Optional[torch.Tensor] = None, ): verify_width = req_to_next_token_ids.shape[1] BLOCK_SIZE = 16 assert verify_width <= BLOCK_SIZE, f"verify_width must be less than {BLOCK_SIZE}" num_reqs = b_req_mtp_start_loc.shape[0] - mtp_step = all_next_token_ids.shape[1] - if req_to_next_token_probs is not None: - assert all_next_token_probs is not None - assert all_next_token_probs.shape == all_next_token_ids.shape + proposal_width = all_next_token_ids.shape[1] + assert proposal_width <= verify_width + if req_to_next_token_scores is not None: + assert schedule_scores is not None + assert schedule_scores.shape == (all_next_token_ids.shape[0], proposal_width - 1) - HAS_NEXT_TOKEN_PROBS = req_to_next_token_probs is not None + HAS_NEXT_TOKEN_SCORES = req_to_next_token_scores is not None # Triton launch arguments cannot be None; static verification uses an unused placeholder. - req_to_next_token_probs_arg = ( - req_to_next_token_probs if req_to_next_token_probs is not None else req_to_next_token_ids - ) - req_to_next_token_probs_stride = ( - req_to_next_token_probs.stride(0) if req_to_next_token_probs is not None else req_to_next_token_ids.stride(0) + req_to_next_token_scores_arg = ( + req_to_next_token_scores if req_to_next_token_scores is not None else req_to_next_token_ids ) - all_next_token_probs_arg = all_next_token_probs if all_next_token_probs is not None else all_next_token_ids - all_next_token_probs_stride = ( - all_next_token_probs.stride(0) if all_next_token_probs is not None else all_next_token_ids.stride(0) + req_to_next_token_scores_stride = ( + req_to_next_token_scores.stride(0) if req_to_next_token_scores is not None else req_to_next_token_ids.stride(0) ) + has_schedule_scores = schedule_scores is not None and schedule_scores.numel() > 0 + schedule_scores_arg = schedule_scores if has_schedule_scores else req_to_next_token_scores_arg + schedule_scores_stride = schedule_scores.stride(0) if has_schedule_scores else 0 grid = (num_reqs,) num_warps = 1 @@ -178,15 +184,16 @@ def mtp_scatter_next_token_ids( req_to_next_token_ids_stride=req_to_next_token_ids.stride(0), all_next_token_ids=all_next_token_ids, all_next_token_ids_stride=all_next_token_ids.stride(0), - req_to_next_token_probs=req_to_next_token_probs_arg, - req_to_next_token_probs_stride=req_to_next_token_probs_stride, - all_next_token_probs=all_next_token_probs_arg, - all_next_token_probs_stride=all_next_token_probs_stride, + req_to_next_token_scores=req_to_next_token_scores_arg, + req_to_next_token_scores_stride=req_to_next_token_scores_stride, + schedule_scores=schedule_scores_arg, + schedule_scores_stride=schedule_scores_stride, spec_accept_len=spec_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, - mtp_step=mtp_step, - HAS_NEXT_TOKEN_PROBS=HAS_NEXT_TOKEN_PROBS, + proposal_width=proposal_width, + verify_width=verify_width, + HAS_NEXT_TOKEN_SCORES=HAS_NEXT_TOKEN_SCORES, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps, num_stages=1, @@ -487,13 +494,13 @@ def prepare_dynamic_spec_model_input( model_input: ModelInput, req_num: int, dynamic_batch_size: int, - req_to_next_token_probs: torch.Tensor, + req_to_next_token_scores: torch.Tensor, pre_draft_step: Optional[int] = None, ): req_num = int(req_num) dynamic_batch_size = int(dynamic_batch_size) assert not model_input.is_prefill, "prepare_dynamic_spec_model_input only supports decode inputs" - assert req_to_next_token_probs is not None + assert req_to_next_token_scores is not None assert dynamic_batch_size >= req_num assert dynamic_batch_size <= model_input.batch_size max_draft_step = int(model_input.draft_step) @@ -508,7 +515,7 @@ def prepare_dynamic_spec_model_input( selected_row_mask = sample_dynamic_spec_row_mask( dynamic_batch_size=dynamic_batch_size, b_req_idx=model_input.b_req_idx, - req_to_next_token_probs=req_to_next_token_probs, + req_to_next_token_scores=req_to_next_token_scores, max_draft_step=max_draft_step, pre_draft_step=pre_draft_step, ) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 4a805cf113..7ff84fd058 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -121,7 +121,7 @@ def __init__(self, max_request_num): dtype=torch.int64, device="cuda", ) - self.req_to_next_token_probs = ( + self.req_to_next_token_scores = ( torch.zeros_like(self.req_to_next_token_ids, dtype=torch.float32) if get_env_start_args().mtp_dynamic_verify else None @@ -143,9 +143,9 @@ def __init__(self, max_request_num): def init_req_sampling_params(self, req: "InferReq"): shm_param = req.sampling_param.shm_param self.req_to_next_token_ids[req.req_idx][0:1].fill_(req.get_last_gen_token()) - if self.req_to_next_token_probs is not None: - self.req_to_next_token_probs[req.req_idx].fill_(0.0) - self.req_to_next_token_probs[req.req_idx][0:1].fill_(1.0) + if self.req_to_next_token_scores is not None: + self.req_to_next_token_scores[req.req_idx].fill_(0.0) + self.req_to_next_token_scores[req.req_idx][0:1].fill_(1.0) self.req_to_presence_penalty[req.req_idx].fill_(shm_param.presence_penalty) self.req_to_frequency_penalty[req.req_idx].fill_(shm_param.frequency_penalty) self.req_to_repetition_penalty[req.req_idx].fill_(shm_param.repetition_penalty) diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py index ba62d247a6..1f7ebe644d 100644 --- a/lightllm/models/qwen3_dflash/model.py +++ b/lightllm/models/qwen3_dflash/model.py @@ -59,7 +59,6 @@ def _init_att_backend(self): raise NotImplementedError("Qwen3 DFlash decode requires FA3") def _init_infer_layer(self, start_layer_index=None): - assert start_layer_index is None self.draft_layer_start = len(self.main_model.layers_infer) self.draft_layer_start += sum( len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models @@ -67,7 +66,6 @@ def _init_infer_layer(self, start_layer_index=None): super()._init_infer_layer(start_layer_index=self.draft_layer_start) def _init_weights(self, start_layer_index=None): - assert start_layer_index is None self.pre_post_weight = self.pre_and_post_weight_class( self.data_type, network_config=self.config, diff --git a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py index 5134fadbd2..2cf407659c 100644 --- a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py +++ b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py @@ -47,7 +47,6 @@ def _markov_step_latent( gate = torch.sigmoid(layer_weight.markov_gate_proj_weight_.mm(gate_input)) return state, gate * prev_embeddings - assert self.markov_head_type_ == "rnn" if state is None: state = torch.zeros_like(prev_embeddings) joint_input = torch.cat([state, prev_embeddings, hidden_states], dim=-1) diff --git a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py index e6c23b0198..437cd36827 100644 --- a/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py +++ b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py @@ -47,7 +47,8 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): weight_names="markov_head.gate_proj.weight", bias_names="markov_head.gate_proj.bias", data_type=self.data_type_, - quant_method=self.quant_cfg.get_quant_method(0, "markov_head.gate_proj"), + # W8A8 MM does not support the Markov projection bias. + quant_method=None, tp_rank=0, tp_world_size=1, ) @@ -58,7 +59,7 @@ def __init__(self, data_type, network_config, quant_cfg: Quantcfg): weight_names="markov_head.joint_proj.weight", bias_names="markov_head.joint_proj.bias", data_type=self.data_type_, - quant_method=self.quant_cfg.get_quant_method(0, "markov_head.joint_proj"), + quant_method=None, tp_rank=0, tp_world_size=1, ) diff --git a/lightllm/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py index 0b33fecfa3..115f817876 100644 --- a/lightllm/models/qwen3_dspark/model.py +++ b/lightllm/models/qwen3_dspark/model.py @@ -20,10 +20,6 @@ class Qwen3DSparkModel(Qwen3DFlashModel): post_layer_infer_class = Qwen3DSparkPostLayerInfer infer_state_class = Qwen3DSparkInferStateInfo - def _verify_params(self): - super()._verify_params() - assert self.config.get("enable_confidence_head", False), "DSpark requires enable_confidence_head=true" - def _token_forward(self, infer_state: Qwen3DSparkInferStateInfo): model_output = super()._token_forward(infer_state) if infer_state.is_cuda_graph: @@ -49,8 +45,6 @@ def _create_unpad_decode_model_output(self, model_output: DSparkModelOutput, ori model_output.draft_token_ids = model_output.draft_token_ids[:origin_batch_size] if model_output.confidence_logits is not None: confidence_rows = model_output.confidence_logits.shape[0] - assert padded_batch_size % confidence_rows == 0 rows_per_confidence = padded_batch_size // confidence_rows - assert origin_batch_size % rows_per_confidence == 0 model_output.confidence_logits = model_output.confidence_logits[: origin_batch_size // rows_per_confidence] return model_output diff --git a/lightllm/models/qwen3_eagle/model.py b/lightllm/models/qwen3_eagle/model.py index 8f54373023..02706f0b28 100644 --- a/lightllm/models/qwen3_eagle/model.py +++ b/lightllm/models/qwen3_eagle/model.py @@ -52,7 +52,6 @@ def _init_mem_manager(self): self.mem_manager = self.main_model.mem_manager def _init_weights(self, start_layer_index=None): - assert start_layer_index is None self.pre_post_weight = self.pre_and_post_weight_class( self.data_type, network_config=self.config, quant_cfg=self.quant_cfg ) @@ -73,7 +72,6 @@ def _init_weights(self, start_layer_index=None): self.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ def _init_infer_layer(self, start_layer_index=None): - assert start_layer_index is None total_pre_layers_num = len(self.main_model.layers_infer) total_pre_layers_num += sum( len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index b4ebd56906..ec7eb82634 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -761,8 +761,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: parser.add_argument( "--mtp_dynamic_verify", action="store_true", - help="""Compatibility switch for dynamic speculative scheduling. - DSpark enables this scheduling mode automatically.""", + help="""Enable dynamic speculative scheduling.""", ) parser.add_argument( "--kv_quant_calibration_config_path", diff --git a/lightllm/server/router/manager.py b/lightllm/server/router/manager.py index 01634d962e..3f750f294f 100644 --- a/lightllm/server/router/manager.py +++ b/lightllm/server/router/manager.py @@ -149,7 +149,7 @@ async def wait_to_model_ready(self): "load_way": self.load_way, "max_total_token_num": self.max_total_token_num, "max_req_num": self.args.running_max_req_size, - "max_seq_length": self.args.max_req_total_len + 8, # 留一点余量 + "max_seq_length": self.args.max_req_total_len + max(8, self.args.mtp_step * 2), "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 6e46b2e75b..91c60f3f9f 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -1,5 +1,4 @@ import os -from collections import Counter import numpy as np import torch @@ -325,7 +324,9 @@ def init_mtp_draft_model(self, main_kvargs: dict): for i in range(draft_model_count): draft_model_cfg, _ = PretrainedConfig.get_config_dict(draft_model_dirs[i]) - if is_chained_draft: + if is_chained_draft or ( + self.enable_decode_microbatch_overlap and spec_mode in ("eagle_with_att", "eagle_no_att") + ): draft_decode_batch_multiplier = self.max_draft_step + 1 elif spec_mode in ("dspark", "dflash"): block_size = int(draft_model_cfg["block_size"]) @@ -814,9 +815,9 @@ def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> def _post_handle( self, run_reqs: List[InferReq], - next_token_ids: List[int], - next_token_logprobs: List[float], - next_token_ranks: List[int], + next_token_ids: torch.Tensor, + next_token_logprobs: torch.Tensor, + next_token_ranks: torch.Tensor, run_reqs_update_packs: List[InferReqUpdatePack], extra_post_req_handle_func: Optional[Callable[[InferReq, int, float], None]] = None, pd_prefill_chunked_handle_func: Optional[Callable[[InferReq, int, float, int], None]] = None, @@ -825,6 +826,10 @@ def _post_handle( extra_post_req_handle_func 用于提供在一个请求确定输出的时候,给出额外的后处理操作,主要是用于 约束输出等模式,设置自己请求内部的状态机的状态,并添加额外的停止判定条件等。 """ + next_token_ids = next_token_ids.tolist() + next_token_logprobs = next_token_logprobs.tolist() + next_token_ranks = next_token_ranks.tolist() + for req_obj, next_token_id, next_token_logprob, next_token_rank, pack in zip( run_reqs, next_token_ids, next_token_logprobs, next_token_ranks, run_reqs_update_packs ): @@ -855,36 +860,6 @@ def _filter_reqs(self, reqs: List[InferReq]): def _trans_req_ids_to_req_objs(self, req_ids: List[int]) -> List[InferReq]: return [g_infer_context.requests_mapping[req_id] for req_id in req_ids] - def _update_spec_accept_ratio( - self, - decode_reqs: List[InferReq], - spec_accept_len_cpu: torch.Tensor, - ): - if self.is_master_in_dp: - for req, accept_len in zip(decode_reqs, spec_accept_len_cpu): - req.update_spec_accepted_token_num(accept_token_num=accept_len - 1) - - return - - def _update_spec_verify_token_num( - self, decode_reqs: List[InferReq], selected_run_reqs: Optional[List[InferReq]] = None - ): - if not self.is_master_in_dp: - return - - if selected_run_reqs is None: - for req in decode_reqs: - req.update_spec_verify_token_num(verify_token_num=req.mtp_step + 1) - req.update_spec_verify_step_num(verify_step_num=1) - return - - verify_rows_by_req = Counter(req.req_idx for req in selected_run_reqs) - for req in decode_reqs: - verify_token_num = verify_rows_by_req[req.req_idx] - if verify_token_num > 0: - req.update_spec_verify_token_num(verify_token_num=verify_token_num) - req.update_spec_verify_step_num(verify_step_num=1) - def _gen_argmax_token_ids(self, model_output: ModelOutput): logits = model_output.logits return torch.argmax(logits, dim=-1) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 50479ee0e3..9476d98ecb 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -2,6 +2,7 @@ import time import torch.distributed as dist from typing import List +from lightllm.common.basemodel.triton_kernel.mtp_utils import gen_b_req_mtp_start_loc from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.mode_backend.overlap_events import OverlapEventPack from lightllm.server.router.model_infer.infer_batch import InferReq @@ -11,6 +12,7 @@ ) from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample from lightllm.server.router.model_infer.infer_batch import g_infer_context +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_current_device_id from .control_state import ControlState @@ -206,7 +208,7 @@ def prefill_mtp( ) # mtp kv fill spec_engine = self.spec_engine - spec_engine.build_initial_draft_state( + spec_engine.build_draft_state_from_prefill( target_model_input=model_input, target_model_output=model_output, next_token_ids=next_token_ids, @@ -245,78 +247,136 @@ def decode_mtp( event_pack: OverlapEventPack, decode_reqs: List[InferReq], ): - """Run the shared speculative draft-and-verify decode flow.""" + """Run the speculative draft-and-verify decode flow.""" model_input, run_reqs = prepare_decode_inputs(decode_reqs) spec_engine = self.spec_engine + req_num = len(decode_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): - spec_plan = spec_engine.plan_decode(model_input=model_input, req_num=len(decode_reqs)) + spec_plan = spec_engine.plan_decode(model_input=model_input, decode_reqs=decode_reqs) - model_input, selected_row_mask = spec_engine.prepare_decode_model_input( + model_input, async_selected_row_mask_cpu = spec_engine.prepare_decode_model_input( model_input=model_input, - req_num=len(decode_reqs), + req_num=req_num, plan=spec_plan, ) - selected_row_mask_cpu = spec_engine.async_copy_selected_row_mask(selected_row_mask) model_output = self.model.forward(model_input) - + if async_selected_row_mask_cpu is not None: + async_selected_row_mask_cpu.wait() + run_reqs = spec_plan.filter_reqs( + reqs=run_reqs, + selected_row_mask_cpu=async_selected_row_mask_cpu.tensor, + ) next_token_ids, next_token_logprobs = sample( model_output.logits, run_reqs, self.eos_id, - selected_row_mask=selected_row_mask, ) next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids) - spec_decode_state = spec_engine.run_decode_speculative_forward( - model_input=model_input, - model_output=model_output, - run_reqs=run_reqs, - req_num=len(decode_reqs), - plan=spec_plan, - selected_row_mask_cpu=selected_row_mask_cpu, + b_req_mtp_start_loc = gen_b_req_mtp_start_loc(model_input.b_mtp_index, num_reqs=req_num) + mtp_accept_len, accepted_index = spec_engine.verify_tokens( + next_token_ids=next_token_ids, + b_req_idx=model_input.b_req_idx, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_mtp_index=model_input.b_mtp_index, + ) + accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="accepted_index", + gpu_tensor=accepted_index, + ) + mtp_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="mtp_accept_len", + gpu_tensor=mtp_accept_len, + ) + + verify_event = torch.cuda.Event() + verify_event.record() + + proposal = spec_engine.propose_next( + main_model_input=model_input, + main_model_output=model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=spec_plan.draft_step, + accept_len=mtp_accept_len, + ) + spec_engine.scatter_next_tokens( + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=proposal.token_ids, + b_req_idx=model_input.b_req_idx, + spec_accept_len=mtp_accept_len, + schedule_scores=proposal.schedule_scores, + ) + + ( + next_token_ids_cpu, + next_token_logprobs_cpu, + next_token_ranks_cpu, + ) = self._async_copy_next_token_infos_to_pin_mem( next_token_ids=next_token_ids, next_token_logprobs=next_token_logprobs, next_token_ranks=next_token_ranks, - copy_next_token_infos=self._async_copy_next_token_infos_to_pin_mem, ) + g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( + b_req_idx=model_input.b_req_idx, + next_token_ids=next_token_ids, + mask=accepted_index == 1, + ) + + sync_event = torch.cuda.Event() + sync_event.record() + # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() - run_reqs, verify_ok_reqs = spec_engine.resolve_decode_pre_post_reqs( - state=spec_decode_state, - decode_reqs=decode_reqs, - ) - self._update_spec_verify_token_num( + run_reqs, verify_ok_reqs = spec_engine.resolve_decode_reqs( + plan=spec_plan, + verify_event=verify_event, + run_reqs=run_reqs, decode_reqs=decode_reqs, - selected_run_reqs=run_reqs if self.enable_dynamic_spec else None, + accepted_index_cpu=accepted_index_cpu, ) + update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) # 第三阶段 event_pack.notify_forward_and_wait_post_handle() - spec_post_state = spec_engine.finish_decode_post( - state=spec_decode_state, - req_num=len(decode_reqs), + sync_event.synchronize() + + spec_engine.update_planner_feedback( + plan=spec_plan, + proposal=proposal, + req_num=req_num, + accept_lengths_cpu=mtp_accept_len_cpu, ) - self._update_spec_accept_ratio( + + spec_engine.record_request_spec_metrics( decode_reqs=decode_reqs, - spec_accept_len_cpu=spec_post_state.spec_accept_len_cpu, + accept_lengths_cpu=mtp_accept_len_cpu, + verified_row_reqs=run_reqs if self.enable_dynamic_spec else None, ) + select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( run_reqs=verify_ok_reqs, - next_token_ids=spec_post_state.next_token_ids, - next_token_logprobs=spec_post_state.next_token_logprobs, - next_token_ranks=spec_post_state.next_token_ranks, + next_token_ids=next_token_ids_cpu[select_mask], + next_token_logprobs=next_token_logprobs_cpu[select_mask], + next_token_ranks=next_token_ranks_cpu[select_mask], run_reqs_update_packs=update_packs, extra_post_req_handle_func=self.extra_post_req_handle_func, ) - if len(spec_post_state.need_free_mem_indexes) > 0: - g_infer_context.req_manager.mem_manager.free(spec_post_state.need_free_mem_indexes) + spec_engine.free_unused_decode_mem( + model_input=model_input, + selected_row_mask_cpu=( + async_selected_row_mask_cpu.tensor if async_selected_row_mask_cpu is not None else None + ), + accepted_index_cpu=accepted_index_cpu, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, + ) # 第四阶段 event_pack.notify_pre_post_handle() diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index e04f70eba3..335ab04a48 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -15,9 +15,6 @@ from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.envs_utils import get_env_start_args from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -from lightllm.common.basemodel.triton_kernel.mtp_utils import ( - linear_att_spec_state_index_update, -) from .control_state import DPControlState @@ -70,19 +67,17 @@ def __init__(self) -> None: @staticmethod def _build_padded_next_token_ids( - token_ids: torch.Tensor, + token_ids: torch.Tensor | None, batch_size: int, copy_len: int, + device: torch.device, source_start: int = 0, ) -> torch.Tensor: """Pad DP-local draft tokens to the collective batch shape.""" copy_len = int(copy_len) source_start = int(source_start) - assert 0 <= copy_len <= int(batch_size) - assert source_start + copy_len <= token_ids.shape[0] - - padded_token_ids = torch.zeros((int(batch_size),), dtype=torch.int64, device=token_ids.device) + padded_token_ids = torch.zeros((int(batch_size),), dtype=torch.int64, device=device) if copy_len > 0: padded_token_ids[:copy_len].copy_( token_ids[source_start : source_start + copy_len], @@ -461,8 +456,9 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] token_ids=next_token_ids, batch_size=model_input.batch_size, copy_len=req_num, + device=model_input.b_req_idx.device, ) - self.spec_engine.build_initial_draft_state( + self.spec_engine.build_draft_state_from_prefill( target_model_input=model_input, target_model_output=model_output, next_token_ids=draft_next_token_ids_gpu, @@ -523,29 +519,21 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): ) = self._async_copy_next_token_infos_to_pin_mem(next_token_ids, next_token_logprobs, next_token_ranks) # verify the next_token_ids - b_req_mtp_start_loc = [index for index, mtp_index in enumerate(b_mtp_index_cpu) if mtp_index == 0] + b_req_mtp_start_loc = [ + index for index, mtp_index in enumerate(b_mtp_index_cpu.tolist()) if mtp_index == 0 + ] b_req_mtp_start_loc = g_pin_mem_manager.gen_from_list( key="b_req_mtp_start_loc", data=b_req_mtp_start_loc, dtype=torch.int32, ).cuda(non_blocking=True) - verify_result = self.spec_engine.verify_target_tokens( - new_next_token_ids=next_token_ids, + spec_accept_len, accepted_index = self.spec_engine.verify_tokens( + next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, + b_mtp_index=model_input.b_mtp_index[0:req_num], ) - spec_accept_len = verify_result.accept_len - accepted_index = verify_result.accepted_index - if self.is_linear_att_mixed_model: - linear_att_spec_state_index_update( - req_to_mtp_state_index=self.model.req_manager.req_to_mtp_state_index, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=b_req_idx, - b_mtp_index=model_input.b_mtp_index[0:req_num], - accepted_index=accepted_index, - verify_width=self.max_draft_step + 1, - ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", gpu_tensor=accepted_index, @@ -580,8 +568,11 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() - self._update_spec_verify_token_num(decode_reqs=decode_reqs) - verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] + self.spec_engine.record_request_spec_metrics( + decode_reqs=decode_reqs, + accept_lengths_cpu=spec_accept_len_cpu, + ) + verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) # 第三阶段 @@ -591,8 +582,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): if eagle_mem_indexes_cpu is not None: need_free_mem_indexes = torch.cat([need_free_mem_indexes, eagle_mem_indexes_cpu], dim=0) - self._update_spec_accept_ratio(decode_reqs=decode_reqs, spec_accept_len_cpu=spec_accept_len_cpu) - select_mask = torch.tensor(accepted_index_cpu, dtype=torch.bool, device="cpu") + select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( run_reqs=verify_ok_reqs, next_token_ids=next_token_ids_cpu[select_mask], @@ -625,11 +615,11 @@ def _draft_decode_vanilla( # share some inference info with the main model draft_model_input = model_input draft_hidden = model_output.spec_hidden - assert draft_hidden is not None draft_next_token_ids_gpu = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input.batch_size, copy_len=req_num, + device=model_input.b_req_idx.device, ) all_next_token_ids.append(draft_next_token_ids_gpu) @@ -641,17 +631,16 @@ def _draft_decode_vanilla( # spec decode: MTP draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) draft_hidden = draft_model_output.spec_hidden - assert draft_hidden is not None draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) all_next_token_ids.append(draft_next_token_ids_gpu) if req_num > 0: - self.spec_engine.scatter_token_id_steps( - token_id_steps=all_next_token_ids, + stacked_next_token_ids = torch.stack(all_next_token_ids, dim=1)[:req_num] + self.spec_engine.scatter_next_tokens( + all_next_token_ids=stacked_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx[:req_num], spec_accept_len=spec_accept_len, - row_count=req_num, ) return None @@ -665,8 +654,6 @@ def _draft_decode_eagle( req_num: int, ): verify_width = self.max_draft_step + 1 - assert req_num % verify_width == 0 - assert model_input.batch_size % verify_width == 0 real_request_num = req_num // verify_width request_capacity = model_input.batch_size // verify_width @@ -674,6 +661,7 @@ def _draft_decode_eagle( token_ids=next_token_ids, batch_size=model_input.batch_size, copy_len=req_num, + device=model_input.b_req_idx.device, ) padded_start_locs = torch.arange( 0, @@ -709,6 +697,7 @@ def _draft_decode_eagle( all_next_token_ids=proposal.token_ids[:req_num], b_req_idx=model_input.b_req_idx[:req_num], spec_accept_len=spec_accept_len, + schedule_scores=(proposal.schedule_scores[:req_num] if proposal.schedule_scores is not None else None), ) return proposal.extra_mem_indexes_cpu @@ -739,6 +728,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I b_mtp_index = torch.cat((model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0) b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) + next_token_ids = torch.empty((0,), dtype=torch.int64, device=logits.device) if (req_num0 + req_num1) > 0: ( next_token_ids, @@ -758,16 +748,18 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I token_ids=next_token_ids, batch_size=model_input0.batch_size, copy_len=req_num0, + device=model_input0.b_req_idx.device, source_start=0, ) draft_next_token_ids_gpu1 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input1.batch_size, copy_len=req_num1, + device=model_input1.b_req_idx.device, source_start=req_num0, ) - self.spec_engine.build_initial_draft_state_overlap( + self.spec_engine.build_draft_state_from_prefill_overlap( target_model_input0=model_input0, target_model_output0=model_output0, next_token_ids0=draft_next_token_ids_gpu0, @@ -842,32 +834,29 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) b_mtp_index_cpu = torch.cat((b_mtp_index_cpu0[0:req_num0], b_mtp_index_cpu1[0:req_num1]), dim=0) - b_req_mtp_start_loc = [index for index, mtp_index in enumerate(b_mtp_index_cpu) if mtp_index == 0] + b_req_mtp_start_loc = [ + index for index, mtp_index in enumerate(b_mtp_index_cpu.tolist()) if mtp_index == 0 + ] b_req_mtp_start_loc = g_pin_mem_manager.gen_from_list( key="b_req_mtp_start_loc", data=b_req_mtp_start_loc, dtype=torch.int32, ).cuda(non_blocking=True) - verify_result = self.spec_engine.verify_target_tokens( - new_next_token_ids=next_token_ids, + b_mtp_index = ( + torch.cat( + (model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), + dim=0, + ) + if self.is_linear_att_mixed_model + else None + ) + spec_accept_len, accepted_index = self.spec_engine.verify_tokens( + next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, + b_mtp_index=b_mtp_index, ) - spec_accept_len = verify_result.accept_len - accepted_index = verify_result.accepted_index - if self.is_linear_att_mixed_model: - b_mtp_index = torch.cat( - (model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0 - ) - linear_att_spec_state_index_update( - req_to_mtp_state_index=self.model.req_manager.req_to_mtp_state_index, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=b_req_idx, - b_mtp_index=b_mtp_index, - accepted_index=accepted_index, - verify_width=self.max_draft_step + 1, - ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", gpu_tensor=accepted_index, @@ -906,8 +895,11 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if req_num0 + req_num1 > 0: event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() - self._update_spec_verify_token_num(decode_reqs=decode_reqs) - verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] + self.spec_engine.record_request_spec_metrics( + decode_reqs=decode_reqs, + accept_lengths_cpu=spec_accept_len_cpu, + ) + verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) event_pack.notify_forward_and_wait_post_handle() @@ -919,8 +911,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if eagle_mem_indexes_cpu is not None: need_free_mem_indexes = torch.cat((need_free_mem_indexes, eagle_mem_indexes_cpu), dim=0) - self._update_spec_accept_ratio(decode_reqs=decode_reqs, spec_accept_len_cpu=spec_accept_len_cpu) - select_mask = torch.tensor(accepted_index_cpu, dtype=torch.bool, device="cpu") + select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( run_reqs=verify_ok_reqs, next_token_ids=next_token_ids_cpu[select_mask], @@ -957,18 +948,19 @@ def _draft_decode_vanilla_overlap( draft_model_input0, draft_model_input1 = model_input0, model_input1 draft_hidden0 = model_output0.spec_hidden draft_hidden1 = model_output1.spec_hidden - assert draft_hidden0 is not None and draft_hidden1 is not None draft_next_token_ids_gpu0 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input0.batch_size, copy_len=req_num0, + device=model_input0.b_req_idx.device, source_start=0, ) draft_next_token_ids_gpu1 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input1.batch_size, copy_len=req_num1, + device=model_input1.b_req_idx.device, source_start=req_num0, ) @@ -984,7 +976,6 @@ def _draft_decode_vanilla_overlap( ) draft_hidden0 = draft_model_output0.spec_hidden draft_hidden1 = draft_model_output1.spec_hidden - assert draft_hidden0 is not None and draft_hidden1 is not None draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) @@ -994,8 +985,9 @@ def _draft_decode_vanilla_overlap( all_next_token_ids.append(draft_next_token_ids) if req_num0 + req_num1 > 0: - self.spec_engine.scatter_token_id_steps( - token_id_steps=all_next_token_ids, + stacked_next_token_ids = torch.stack(all_next_token_ids, dim=1) + self.spec_engine.scatter_next_tokens( + all_next_token_ids=stacked_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, spec_accept_len=spec_accept_len, @@ -1016,10 +1008,6 @@ def _draft_decode_eagle_overlap( req_num1: int = 0, ): verify_width = self.max_draft_step + 1 - assert req_num0 % verify_width == 0 - assert req_num1 % verify_width == 0 - assert model_input0.batch_size % verify_width == 0 - assert model_input1.batch_size % verify_width == 0 real_request_num0 = req_num0 // verify_width real_request_num1 = req_num1 // verify_width request_capacity0 = model_input0.batch_size // verify_width @@ -1029,12 +1017,14 @@ def _draft_decode_eagle_overlap( token_ids=next_token_ids, batch_size=model_input0.batch_size, copy_len=req_num0, + device=model_input0.b_req_idx.device, source_start=0, ) padded_next_token_ids1 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input1.batch_size, copy_len=req_num1, + device=model_input1.b_req_idx.device, source_start=req_num0, ) padded_accept_len0 = torch.ones( @@ -1074,5 +1064,6 @@ def _draft_decode_eagle_overlap( all_next_token_ids=proposal.token_ids, b_req_idx=b_req_idx, spec_accept_len=spec_accept_len, + schedule_scores=proposal.schedule_scores, ) return proposal.extra_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py index 1157275afb..5b29ea0510 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_post_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_post_process.py @@ -1,6 +1,5 @@ import torch -from typing import List, Tuple, Optional -from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import trim_post_sample_tensors +from typing import List, Tuple from lightllm.common.basemodel.triton_kernel.post_process.apply_penalty import apply_penalty from lightllm.common.basemodel.triton_kernel.post_process.apply_penalty_gpu_cache import apply_penalty_gpu_cache from lightllm.common.basemodel.triton_kernel.post_process.apply_invalid_token import apply_invalid_token_ids @@ -9,12 +8,7 @@ from lightllm.utils.envs_utils import get_env_start_args -def sample( - logits: torch.Tensor, - reqs: List[InferReq], - eos_id: List[int] = [2], - selected_row_mask: Optional[torch.Tensor] = None, -): +def sample(logits: torch.Tensor, reqs: List[InferReq], eos_id: List[int] = [2]): ( b_req_idx, b_temperatures, @@ -33,34 +27,6 @@ def sample( eos_ids = g_pin_mem_manager.gen_from_list(key="eos_ids", data=eos_id, dtype=torch.int32).cuda(non_blocking=True) sampling_params_manager = g_infer_context.req_manager.req_sampling_params_manager - if selected_row_mask is not None: - ( - b_req_idx, - b_temperatures, - b_top_ps, - b_top_ks, - b_length_penalty_param, - b_mask_eos_reqs, - ) = trim_post_sample_tensors( - dynamic_batch_size=logits.shape[0], - selected_row_mask=selected_row_mask, - b_req_idx=b_req_idx, - b_temperatures=b_temperatures, - b_top_ps=b_top_ps, - b_top_ks=b_top_ks, - b_length_penalty_param=b_length_penalty_param, - b_mask_eos_reqs=b_mask_eos_reqs, - ) - if ( - has_invalid_token_ids - or exist_req_use_random_seed - or sampling_params_manager.penalty_counter_mode == "cpu_counter" - ): - reqs = _get_selected_reqs(reqs=reqs, selected_row_mask=selected_row_mask) - if has_invalid_token_ids: - invalid_token_ids, cu_invalid_token_num, has_invalid_token_ids = _get_invalid_token_tensors(reqs=reqs) - if exist_req_use_random_seed: - exist_req_use_random_seed = any(req.generator is not None for req in reqs) # 这里需要区分历史token的频率惩罚类的系数的生效模式,目前支持两种在线统计方式: # 一种是基于 cpu 的,每个 req 对象利用其上绑定的dict对象out_token_id_count,每生成一个token就进行相应 @@ -265,37 +231,3 @@ def _get_post_sample_tensors(reqs: List[InferReq]): skip_top_p, exist_req_use_random_seed, ) - - -def _get_selected_reqs(reqs: List[InferReq], selected_row_mask: torch.Tensor): - selected_row_mask_cpu = selected_row_mask.detach().cpu().tolist() - return [req_obj for req_obj, selected in zip(reqs, selected_row_mask_cpu) if int(selected) != 0] - - -def _get_invalid_token_tensors(reqs: List[InferReq]): - invalid_token_ids: List[int] = [] - has_invalid_token_ids = False - cu_invalid_token_num = [0] - invalid_token_num_start = 0 - - for req_obj in reqs: - invalid_token_num_start += len(req_obj.sampling_param.invalid_token_ids) - cu_invalid_token_num.append(invalid_token_num_start) - if len(req_obj.sampling_param.invalid_token_ids) > 0: - has_invalid_token_ids = True - invalid_token_ids.extend(req_obj.sampling_param.invalid_token_ids) - - if not has_invalid_token_ids: - return None, None, False - - invalid_token_ids_cpu = g_pin_mem_manager.gen_from_list( - key="invalid_token_ids", data=invalid_token_ids, dtype=torch.int32 - ) - cu_invalid_token_num_cpu = g_pin_mem_manager.gen_from_list( - key="cu_invalid_token_num", data=cu_invalid_token_num, dtype=torch.int32 - ) - return ( - invalid_token_ids_cpu.cuda(non_blocking=True), - cu_invalid_token_num_cpu.cuda(non_blocking=True), - True, - ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 8daa4651d9..911929ef89 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -229,10 +229,11 @@ def build_spec_shared_group_markers(b_mtp_index: torch.Tensor) -> torch.Tensor: # Each logical request starts at row index 0. Only the final row of each # speculative query group stores its group size; earlier rows store zero. max_batch_shared_group_size = get_diverse_max_batch_shared_group_size() + mtp_indexes = b_mtp_index.tolist() b_mark_shared_group = [] group_start = 0 - for group_end in range(1, len(b_mtp_index) + 1): - reaches_request_boundary = group_end == len(b_mtp_index) or b_mtp_index[group_end] == 0 + for group_end in range(1, len(mtp_indexes) + 1): + reaches_request_boundary = group_end == len(mtp_indexes) or mtp_indexes[group_end] == 0 reaches_size_limit = group_end - group_start == max_batch_shared_group_size if not reaches_request_boundary and not reaches_size_limit: continue @@ -242,6 +243,6 @@ def build_spec_shared_group_markers(b_mtp_index: torch.Tensor) -> torch.Tensor: b_mark_shared_group.append(group_size) group_start = group_end - assert len(b_mark_shared_group) == len(b_mtp_index) + assert len(b_mark_shared_group) == len(mtp_indexes) b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cpu") return b_mark_shared_group diff --git a/lightllm/server/router/model_infer/pin_mem_manager.py b/lightllm/server/router/model_infer/pin_mem_manager.py index d5553e4498..7ca218bcd8 100644 --- a/lightllm/server/router/model_infer/pin_mem_manager.py +++ b/lightllm/server/router/model_infer/pin_mem_manager.py @@ -1,9 +1,19 @@ import torch import threading import collections +from dataclasses import dataclass from typing import List, Dict, Union, Sequence +@dataclass +class AsyncPinnedCpuTensor: + tensor: torch.Tensor + ready_event: torch.cuda.Event + + def wait(self) -> None: + self.ready_event.synchronize() + + class PinMemTensorManager: def __init__(self): self.lock = threading.Lock() @@ -49,6 +59,16 @@ def async_copy_from_gpu_tensor(self, key: str, gpu_tensor: torch.Tensor) -> torc pin_mem.copy_(gpu_tensor.view(-1), non_blocking=True) return pin_mem.view(gpu_tensor.shape) + def async_copy_from_gpu_tensor_with_event( + self, + key: str, + gpu_tensor: torch.Tensor, + ) -> AsyncPinnedCpuTensor: + cpu_tensor = self.async_copy_from_gpu_tensor(key=key, gpu_tensor=gpu_tensor) + ready_event = torch.cuda.Event() + ready_event.record() + return AsyncPinnedCpuTensor(tensor=cpu_tensor, ready_event=ready_event) + def get_const_cpu_tensor( self, key: str, diff --git a/lightllm/server/router/model_infer/speculative/engine.py b/lightllm/server/router/model_infer/speculative/engine.py index 0ebe83bd8e..c8a70f0f0a 100644 --- a/lightllm/server/router/model_infer/speculative/engine.py +++ b/lightllm/server/router/model_infer/speculative/engine.py @@ -1,81 +1,62 @@ from __future__ import annotations -from typing import Callable, List, Optional, Tuple +from collections import Counter +from typing import List, Optional, Tuple import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.mtp_utils import ( + linear_att_spec_state_index_update, mtp_scatter_next_token_ids, mtp_verify, ) -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager from lightllm.server.router.model_infer.speculative.planner import ( - DSparkDynamicSpecPlanner, - DynamicSpecPlanner, - Eagle3DynamicSpecPlanner, + DSparkPlanner, FixedSpecPlanner, + LightSpecPlanner, SpecDecodePlan, ) from lightllm.server.router.model_infer.speculative.proposers import build_spec_proposer from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.speculative.runner import ( - SpecDecodeForwardState, - SpecDecodePostState, - SpecDecodeRunner, -) class SpecEngine: - """Coordinates speculative decoding for one model backend. - - The engine handles the common decode flow, while each proposer owns its - algorithm-specific draft-state update and proposal generation. + """Owns speculative planning, verification, and proposal generation. - Each decode iteration: - 1. the target model samples tokens and returns captured hidden states; - 2. verification computes the accepted prefix against the previous proposal; - 3. the proposer updates draft state and generates the next proposal; - 4. the proposal is stored in per-request buffers for the next iteration. + The model backend controls target forward, sampling, stream synchronization, + and request post-processing. Each proposer owns its algorithm-specific draft + state and proposal generation. """ def __init__(self, backend, spec_mode: str, enable_dynamic_spec: bool) -> None: self.backend = backend self.spec_mode = spec_mode self.enable_dynamic_spec = enable_dynamic_spec - self.proposer = build_spec_proposer(engine=self) - self.decode_runner = SpecDecodeRunner(engine=self) + self.proposer = build_spec_proposer( + spec_mode=spec_mode, + backend=backend, + enable_dynamic_spec=enable_dynamic_spec, + ) self.planner = self._build_decode_planner() self._register_cuda_graph_costs() - self._dynamic_accept_stats_calls = 0 - - def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: - """Allocate speculative draft-owned temporary KV slots.""" - - token_count = int(token_count) - assert token_count >= 0 - if token_count == 0: - return torch.empty((0,), dtype=torch.int32, device="cpu") - from lightllm.server.router.model_infer.infer_batch import g_infer_context + # Prefill draft-state initialization. - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) - return g_infer_context.req_manager.mem_manager.alloc(token_count) - - def build_initial_draft_state( + def build_draft_state_from_prefill( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - self.proposer.build_initial_draft_state( + self.proposer.build_draft_state_from_prefill( target_model_input=target_model_input, target_model_output=target_model_output, next_token_ids=next_token_ids, ) - def build_initial_draft_state_overlap( + def build_draft_state_from_prefill_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -84,7 +65,7 @@ def build_initial_draft_state_overlap( target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - self.proposer.build_initial_draft_state_overlap( + self.proposer.build_draft_state_from_prefill_overlap( target_model_input0=target_model_input0, target_model_output0=target_model_output0, next_token_ids0=next_token_ids0, @@ -93,56 +74,30 @@ def build_initial_draft_state_overlap( next_token_ids1=next_token_ids1, ) - def plan_decode(self, model_input: ModelInput, req_num: int) -> SpecDecodePlan: + # Decode planning and target verification. + + def plan_decode(self, model_input: ModelInput, decode_reqs: List) -> SpecDecodePlan: """Return the fixed or dynamic speculative plan for one decode iteration.""" + req_num = len(decode_reqs) + if isinstance(self.planner, LightSpecPlanner): + # Prefill initializes draft cache but does not create a proposal. + # After one decode, cur_output_len > 1 and the request owns the + # previous iteration's candidates. + proposal_req_num = sum(req.cur_output_len > 1 for req in decode_reqs) + return self.planner.plan( + req_num=req_num, + original_batch_size=model_input.batch_size, + proposal_req_num=proposal_req_num, + ) return self.planner.plan(req_num=req_num, original_batch_size=model_input.batch_size) - def run_decode_speculative_forward( - self, - model_input: ModelInput, - model_output: ModelOutput, - run_reqs: List, - req_num: int, - plan: SpecDecodePlan, - selected_row_mask_cpu: Optional[torch.Tensor], - next_token_ids: torch.Tensor, - next_token_logprobs: torch.Tensor, - next_token_ranks: torch.Tensor, - copy_next_token_infos: Callable[ - [torch.Tensor, torch.Tensor, torch.Tensor], - Tuple[torch.Tensor, torch.Tensor, torch.Tensor], - ], - ) -> SpecDecodeForwardState: - return self.decode_runner.run_speculative_forward( - model_input=model_input, - model_output=model_output, - run_reqs=run_reqs, - req_num=req_num, - plan=plan, - selected_row_mask_cpu=selected_row_mask_cpu, - next_token_ids=next_token_ids, - next_token_logprobs=next_token_logprobs, - next_token_ranks=next_token_ranks, - copy_next_token_infos=copy_next_token_infos, - ) - - def resolve_decode_pre_post_reqs(self, state: SpecDecodeForwardState, decode_reqs: List): - return self.decode_runner.resolve_pre_post_reqs(state=state, decode_reqs=decode_reqs) - - def finish_decode_post( - self, - state: SpecDecodeForwardState, - req_num: int, - ) -> SpecDecodePostState: - return self.decode_runner.finish_post(state=state, req_num=req_num) - def prepare_decode_model_input( self, model_input: ModelInput, req_num: int, plan: SpecDecodePlan, - ): + ) -> Tuple[ModelInput, Optional[AsyncPinnedCpuTensor]]: """Apply target verify-row compaction when the dynamic planner selects it.""" if not plan.is_dynamic: @@ -151,143 +106,66 @@ def prepare_decode_model_input( if plan.dynamic_batch_size == model_input.batch_size: return model_input, None - self._clear_stale_dynamic_token_probs(pre_draft_step=plan.pre_draft_step) - from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_spec_model_input model_input, selected_row_mask = prepare_dynamic_spec_model_input( model_input=model_input, req_num=req_num, dynamic_batch_size=plan.dynamic_batch_size, - req_to_next_token_probs=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs, + req_to_next_token_scores=( + self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_scores + ), pre_draft_step=plan.pre_draft_step, ) - return model_input, selected_row_mask - - def async_copy_selected_row_mask(self, selected_row_mask: Optional[torch.Tensor]): - if selected_row_mask is None: - return None - return g_pin_mem_manager.async_copy_from_gpu_tensor( + selected_row_mask_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor_with_event( key="selected_row_mask", gpu_tensor=selected_row_mask, ) + return model_input, selected_row_mask_cpu - def build_decode_req_lists( + def verify_tokens( self, - original_run_reqs, - selected_row_mask_cpu: Optional[torch.Tensor], - accepted_index_cpu: torch.Tensor, - ): - """Build post-handle request lists after optional verify-row compaction.""" - - if selected_row_mask_cpu is not None: - run_reqs = [req for req, selected in zip(original_run_reqs, selected_row_mask_cpu.tolist()) if selected] - else: - run_reqs = original_run_reqs - - verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] - return run_reqs, verify_ok_reqs + next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + b_mtp_index: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + accept_lengths, accepted_index = mtp_verify( + req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + new_next_token_ids=next_token_ids, + b_req_idx=b_req_idx, + ) + if self.backend.is_linear_att_mixed_model: + assert b_mtp_index is not None + linear_att_spec_state_index_update( + req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_req_idx=b_req_idx, + b_mtp_index=b_mtp_index, + accepted_index=accepted_index, + verify_width=self.backend.max_draft_step + 1, + ) + return accept_lengths, accepted_index - def build_decode_free_mem_indexes_cpu( + def resolve_decode_reqs( self, - model_input: ModelInput, - selected_row_mask_cpu: Optional[torch.Tensor], - accepted_index_cpu: torch.Tensor, - ) -> torch.Tensor: - mem_indexes_cpu = model_input.mem_indexes_cpu - if selected_row_mask_cpu is None: - return mem_indexes_cpu[accepted_index_cpu == 0] - - selected_mask = selected_row_mask_cpu.to(dtype=torch.bool) - accepted_mask = accepted_index_cpu.to(dtype=torch.bool) - selected_mem_indexes_cpu = mem_indexes_cpu[selected_mask] - assert selected_mem_indexes_cpu.shape[0] == accepted_mask.shape[0] - - unselected_mem_indexes_cpu = mem_indexes_cpu[~selected_mask] - rejected_selected_mem_indexes_cpu = selected_mem_indexes_cpu[~accepted_mask] - if len(unselected_mem_indexes_cpu) == 0: - return rejected_selected_mem_indexes_cpu - if len(rejected_selected_mem_indexes_cpu) == 0: - return unselected_mem_indexes_cpu - return torch.cat([unselected_mem_indexes_cpu, rejected_selected_mem_indexes_cpu], dim=0) - - def update_dynamic_accept_stats( - self, - req_num: int, + plan: SpecDecodePlan, + verify_event: torch.cuda.Event, + run_reqs: List, + decode_reqs: List, accepted_index_cpu: torch.Tensor, - spec_accept_len_cpu: torch.Tensor, - dynamic_batch_size: int, - pre_draft_step: int, - ) -> None: - assert accepted_index_cpu.shape[0] == dynamic_batch_size - assert spec_accept_len_cpu.shape[0] == req_num - accept_lengths = spec_accept_len_cpu.numpy() - accept_count = int(accept_lengths.sum()) - total_count = int(dynamic_batch_size) - is_full_verify = dynamic_batch_size == req_num * (self.backend.max_draft_step + 1) - - # The first target decode has no preceding draft proposal, so every - # request structurally accepts only its base token. Treating that - # cold-start iteration as a K-wide acceptance sample makes a highly - # predictable workload look maximally hard and can collapse the - # controller before the first real proposal is verified. - self._dynamic_accept_stats_calls += 1 - if self._dynamic_accept_stats_calls == 1: - return + ) -> Tuple[List, List]: + """Return requests ready for pre/post handling after verification.""" - if self.spec_mode == "eagle3": - self.planner.update_req_num_to_dynamic_batch_size_to_accept_ratio( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - accept_ratio=accept_count / total_count, - pre_draft_step=pre_draft_step, - ) - if is_full_verify: - self.planner.update_full_verify_tokens_per_req( - accept_count / req_num, - req_num=req_num, - ) - self.planner.update_observed_iteration_stats( - tokens_per_req=accept_count / req_num, - verify_rows_per_req=dynamic_batch_size / req_num, - is_full_verify=is_full_verify, - req_num=req_num, - ) - # Confidence-selected rows bias the survival curve upward, so only - # full-width verification provides valid Eagle3 prefix samples. - if is_full_verify: - self.planner.update_verified_batch_prefix_stats( - verify_and_accept_lengths=[ - (self.backend.max_draft_step + 1, int(accept_len)) for accept_len in accept_lengths - ], - ) - else: - self.planner.update_req_num_to_dynamic_batch_size_to_accept_ratio( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - accept_ratio=accept_count / total_count, - ) - verify_rows_per_req = max(1, int(round(dynamic_batch_size / req_num))) - for accept_len in accept_lengths: - self.planner.update_verified_prefix_stats( - verify_len=verify_rows_per_req, - accept_len=int(accept_len), - ) + if plan.skip_verify_sync: + return decode_reqs, decode_reqs - def needs_schedule_probs_cpu(self) -> bool: - """Whether this planner consumes proposal confidence on the CPU.""" - - return self.enable_dynamic_spec and self.spec_mode == "dspark" + verify_event.synchronize() + verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] + return run_reqs, verify_ok_reqs - def update_dynamic_schedule_stats( - self, - req_num: int, - schedule_probs_cpu: torch.Tensor, - ) -> None: - self.planner.update_predicted_schedule_probs( - schedule_probs=schedule_probs_cpu, - req_num=req_num, - ) + # Draft proposal generation and persistence. def propose_next( self, @@ -321,6 +199,9 @@ def propose_next_overlap( accept_len1: torch.Tensor, draft_step: int, ) -> SpecProposal: + # TODO: DP overlap currently disables dynamic drafting and always uses + # max_draft_step. Share the score-feedback path with propose_next when + # dynamic overlap is supported. return self.proposer.propose_next_overlap( main_model_input0=main_model_input0, main_model_output0=main_model_output0, @@ -335,77 +216,13 @@ def propose_next_overlap( draft_step=draft_step, ) - def verify_target_tokens( - self, - new_next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - ) -> Tuple[torch.Tensor, torch.Tensor]: - return mtp_verify( - req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - new_next_token_ids=new_next_token_ids, - b_req_idx=b_req_idx, - ) - - def build_all_next_token_probs( - self, - proposal: SpecProposal, - draft_step: int, - ) -> Optional[torch.Tensor]: - """Build selected-token probabilities for dynamic speculative scatter. - - Output shape is [verify_batch, max_draft_step + 1]. Column 0 is the target - token probability, fixed to 1 because the target sample is always the - base accepted position. Draft columns store selected-token - probabilities from each proposer step. - """ - - if not self.enable_dynamic_spec: - return None - - schedule_probs = proposal.schedule_probs if proposal.schedule_probs is not None else proposal.draft_probs - assert schedule_probs is not None - - verify_row_count = proposal.token_ids.shape[0] - all_next_token_probs = torch.zeros( - size=(verify_row_count, self.backend.max_draft_step + 1), - dtype=torch.float32, - device=proposal.token_ids.device, - ) - all_next_token_probs[:, 0] = 1.0 - - if isinstance(schedule_probs, torch.Tensor): - assert schedule_probs.shape == (verify_row_count, draft_step) - if draft_step > 0: - all_next_token_probs[:, 1 : draft_step + 1] = schedule_probs - return all_next_token_probs - - assert len(schedule_probs) == draft_step - for step_idx, step_probs in enumerate(schedule_probs): - all_next_token_probs[:, step_idx + 1] = step_probs - return all_next_token_probs - - def pad_all_next_token_ids(self, token_ids: torch.Tensor, draft_step: int) -> torch.Tensor: - """Pad a shorter dynamic proposal to the configured speculative width.""" - - if not self.enable_dynamic_spec or draft_step >= self.backend.max_draft_step: - return token_ids - - append_next_token_ids = torch.ones( - size=(token_ids.shape[0], self.backend.max_draft_step - draft_step), - dtype=token_ids.dtype, - device=token_ids.device, - ) - return torch.cat([token_ids, append_next_token_ids], dim=-1) - def scatter_next_tokens( self, b_req_mtp_start_loc: torch.Tensor, all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, spec_accept_len: torch.Tensor, - all_next_token_probs: Optional[torch.Tensor] = None, + schedule_scores: Optional[torch.Tensor] = None, ) -> None: mtp_scatter_next_token_ids( req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, @@ -413,43 +230,94 @@ def scatter_next_tokens( all_next_token_ids=all_next_token_ids, b_req_idx=b_req_idx, spec_accept_len=spec_accept_len, - req_to_next_token_probs=( - self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs - if all_next_token_probs is not None + req_to_next_token_scores=( + self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_scores + if schedule_scores is not None else None ), - all_next_token_probs=all_next_token_probs, + schedule_scores=schedule_scores, ) - def scatter_token_id_steps( + # Planner feedback, request metrics, and resource cleanup. + + def update_planner_feedback( self, - token_id_steps: List[torch.Tensor], - b_req_mtp_start_loc: torch.Tensor, - b_req_idx: torch.Tensor, - spec_accept_len: torch.Tensor, - row_count: Optional[int] = None, - ) -> torch.Tensor: - """Stack proposal token columns and scatter them for the next verify.""" - - all_next_token_ids = torch.stack(token_id_steps, dim=1) - if row_count is not None: - all_next_token_ids = all_next_token_ids[: int(row_count), :] - self.scatter_next_tokens( - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=b_req_idx, - spec_accept_len=spec_accept_len, + plan: SpecDecodePlan, + proposal: SpecProposal, + req_num: int, + accept_lengths_cpu: torch.Tensor, + ) -> None: + """Feed iteration-level observations into the dynamic planner.""" + + if not self.enable_dynamic_spec: + return + + self.planner.update_feedback( + plan=plan, + req_num=req_num, + accept_lengths=accept_lengths_cpu, + schedule_scores=proposal.schedule_scores_cpu, ) - return all_next_token_ids + + def record_request_spec_metrics( + self, + decode_reqs: List, + accept_lengths_cpu: torch.Tensor, + verified_row_reqs: Optional[List] = None, + ) -> None: + """Accumulate user-visible speculative metrics on each request.""" + + if not self.backend.is_master_in_dp: + return + + accept_lengths = accept_lengths_cpu.tolist() + assert len(accept_lengths) == len(decode_reqs) + verify_rows_by_req = None if verified_row_reqs is None else Counter(req.req_idx for req in verified_row_reqs) + for req, accept_len in zip(decode_reqs, accept_lengths): + req.update_spec_accepted_token_num(accept_token_num=accept_len - 1) + verify_token_num = req.mtp_step + 1 if verify_rows_by_req is None else verify_rows_by_req[req.req_idx] + if verify_token_num > 0: + req.update_spec_verify_token_num(verify_token_num=verify_token_num) + req.update_spec_verify_step_num(verify_step_num=1) + + def free_unused_decode_mem( + self, + model_input: ModelInput, + selected_row_mask_cpu: Optional[torch.Tensor], + accepted_index_cpu: torch.Tensor, + extra_mem_indexes_cpu: Optional[torch.Tensor], + ) -> None: + """Free rejected target KV slots and draft-only temporary slots.""" + + mem_indexes_cpu = model_input.mem_indexes_cpu + if selected_row_mask_cpu is None: + free_mask = accepted_index_cpu == 0 + else: + selected_mask = selected_row_mask_cpu.to(dtype=torch.bool) + free_mask = selected_mask.logical_not() + free_mask[selected_mask] = accepted_index_cpu == 0 + need_free_mem_indexes = mem_indexes_cpu[free_mask] + + if extra_mem_indexes_cpu is not None: + need_free_mem_indexes = torch.cat([need_free_mem_indexes, extra_mem_indexes_cpu], dim=0) + if len(need_free_mem_indexes) > 0: + self.backend.model.req_manager.mem_manager.free(need_free_mem_indexes) + + # Construction helpers. def _build_decode_planner(self): if not self.enable_dynamic_spec: return FixedSpecPlanner(max_draft_step=self.backend.max_draft_step) if self.spec_mode == "dspark": - return DSparkDynamicSpecPlanner(max_draft_step=self.backend.max_draft_step) - if self.spec_mode == "eagle3": - return Eagle3DynamicSpecPlanner(max_draft_step=self.backend.max_draft_step) - return DynamicSpecPlanner(max_draft_step=self.backend.max_draft_step) + return DSparkPlanner( + max_draft_step=self.backend.max_draft_step, + draft_cost_provider=self.proposer, + ) + + return LightSpecPlanner( + max_draft_step=self.backend.max_draft_step, + draft_cost_provider=self.proposer, + ) def _register_cuda_graph_costs(self) -> None: if not self.enable_dynamic_spec: @@ -464,8 +332,6 @@ def _register_cuda_graph_costs(self) -> None: is_draft_model=False, ) - if self.spec_mode == "dspark": - return for draft_model in self.backend.draft_models: draft_graph = draft_model.graph if draft_graph is None: @@ -476,10 +342,3 @@ def _register_cuda_graph_costs(self) -> None: infer_cost_ms=infer_cost_ms, is_draft_model=True, ) - - def _clear_stale_dynamic_token_probs(self, pre_draft_step: int) -> None: - # Columns after the previous draft length are stale and must not be - # sampled by dynamic row compaction in the current target forward. - self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_probs[ - :, (pre_draft_step + 1) : - ].fill_(0.0) diff --git a/lightllm/server/router/model_infer/speculative/planner.py b/lightllm/server/router/model_infer/speculative/planner.py index 33bb5f737d..d8c954caa6 100644 --- a/lightllm/server/router/model_infer/speculative/planner.py +++ b/lightllm/server/router/model_infer/speculative/planner.py @@ -1,15 +1,15 @@ from __future__ import annotations -import math -import os -import random from collections import deque from dataclasses import dataclass -from typing import Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Dict, List, Optional, Tuple import numpy as np from sortedcontainers import SortedDict +if TYPE_CHECKING: + from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer + @dataclass(frozen=True) class SpecDecodePlan: @@ -29,6 +29,10 @@ class SpecDecodePlan: dynamic_batch_size: Optional[int] draft_step: int pre_draft_step: int + # False when the current batch contains requests without a proposal from + # the previous iteration. Such an iteration does not represent one + # well-defined LightSpec runtime configuration. + record_progress: bool = True @property def is_dynamic(self) -> bool: @@ -38,6 +42,9 @@ def is_dynamic(self) -> bool: def skip_verify_sync(self) -> bool: return self.is_dynamic and self.pre_draft_step == 0 + def filter_reqs(self, reqs: List, selected_row_mask_cpu) -> List: + return [req for req, selected in zip(reqs, selected_row_mask_cpu.tolist()) if selected] + class FixedSpecPlanner: """Planner for fixed-width speculative decoding.""" @@ -53,926 +60,318 @@ def plan(self, req_num: int | None = None, original_batch_size: int | None = Non ) -class DynamicSpecPlanner: +class LightSpecPlanner: + """Choose the current verify budget and the next draft configuration. + + LightSpec evaluates a runtime configuration as ``(N, B, d)``: logical + requests, physical target rows, and draft configuration. Planning follows + the paper's budget-then-fill order. It first chooses ``B`` for the proposal + produced by ``pre_draft_step``, then keeps that ``B`` fixed while choosing + the ``draft_step`` that will produce the next proposal. Candidate identities + and per-request verify widths belong to the GPU Fill stage, not this planner. + The next draft depth bounds the next iteration's verify budget; it does not + need to reproduce the verify budget consumed in the current iteration. + + Each configuration minimizes estimated milliseconds per committed token: + ``(target_cost(B) + draft_cost(N, B, d)) / expected_progress(N, B, d)``. + Progress is learned from one normalized batch observation ``sum(accept_len) + / B`` per iteration, then smoothed with an EMA. It is never updated once per + request; doing so would make adaptation depend on concurrency. + + The proposer supplies its valid draft configurations and complete draft + cost. The planner therefore stays independent of the proposal algorithm and + its physical execution layout. + """ + def __init__( self, max_draft_step: int, - use_random_mode: bool = True, - random_mode_iter_threshold: int = 100, + draft_cost_provider: "BaseSpecProposer", ) -> None: self.max_draft_step = int(max_draft_step) + self.draft_cost_provider = draft_cost_provider + self.draft_steps = draft_cost_provider.get_draft_steps() - # 用于记录 decode 时的静态推理耗时(ms)。 self.target_infer_costs = _InferCostMsTable() self.draft_infer_costs = _InferCostMsTable() - # 记录不同 draft 深度的接受概率;第一个 target token 必然接受,不需要统计。 - self.draft_len_to_accept_ratio = [ - _EMAValue(decay=0.95, init_value=1.0, enable_decay_warmup=False) for _ in range(self.max_draft_step) - ] - # 记录请求数量以及对应的推理dynamic_batch_size 对应的接受率统计 - self.req_num_to_dynamic_batch_size_to_accept_ratio: Dict[int, Dict[int, _EMAValue]] = {} - - # 每多少个请求采用随机的方式决定 dynamic_batch_size - self._iter = 0 - self._iter_threshold = int(random_mode_iter_threshold) - self._use_random_mode = bool(use_random_mode) - self._random = random.Random(0) + # Each observation is the normalized committed progress U / B from one + # complete batch. The draft configuration is part of the key because a + # deeper candidate pool may improve Fill even at the same (N, B). + self.progress_ema_by_config: Dict[Tuple[int, int, int], _EMAValue] = {} - # 记录上一次选择的draft step 步长,才好选择对应的 dynamic_batch_size + # The current verify width is bounded by the proposal built last time. self.pre_draft_step = self.max_draft_step - def plan(self, req_num: int, original_batch_size: int) -> SpecDecodePlan: - dynamic_batch_size, draft_step, pre_draft_step = self.get_dynamic_batch_size( - req_num=req_num, - original_batch_size=original_batch_size, - ) - return SpecDecodePlan( - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - pre_draft_step=pre_draft_step, - ) - - def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: - cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs - cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - - def update_draft_len_to_accept_ratio(self, draft_len: int, accept_ratio: float) -> None: - assert draft_len > 0 and draft_len <= self.max_draft_step - self.draft_len_to_accept_ratio[draft_len - 1].update(accept_ratio) - - def update_verified_prefix_stats(self, verify_len: int, accept_len: int) -> None: - if verify_len - 1 <= 0: - return - for draft_index in range(verify_len - 1): - draft_len = draft_index + 1 - ratio = (accept_len - 1) / draft_len - ratio = max(0.0, min(1.0, ratio)) - self.update_draft_len_to_accept_ratio( - draft_len=draft_len, - accept_ratio=ratio, - ) - - def update_req_num_to_dynamic_batch_size_to_accept_ratio( - self, req_num: int, dynamic_batch_size: int, accept_ratio: float - ) -> None: - assert dynamic_batch_size >= req_num - self._get_req_num_to_dynamic_batch_size_to_accept_ratio( - req_num=req_num, dynamic_batch_size=dynamic_batch_size - ).update(accept_ratio) - - def _get_req_num_to_dynamic_batch_size_to_accept_ratio(self, req_num: int, dynamic_batch_size: int) -> "_EMAValue": - assert dynamic_batch_size >= req_num - if req_num not in self.req_num_to_dynamic_batch_size_to_accept_ratio: - self.req_num_to_dynamic_batch_size_to_accept_ratio[req_num] = {} - if dynamic_batch_size not in self.req_num_to_dynamic_batch_size_to_accept_ratio[req_num]: - self.req_num_to_dynamic_batch_size_to_accept_ratio[req_num][dynamic_batch_size] = _EMAValue( - decay=0.9, init_value=1.0, enable_decay_warmup=True - ) - return self.req_num_to_dynamic_batch_size_to_accept_ratio[req_num][dynamic_batch_size] - - def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: - """ - 返回 (dynamic_batch_size, draft_step, pre_draft_step)。 - pre_draft_step 是上一轮推理实际使用的 draft_step, - 调用方可据此判断当前 verify 是否有真实候选需要验证 - (pre_draft_step == 0 时 accept_len 恒为 1,无需等待 GPU verify 结果)。 - """ - assert req_num * (self.max_draft_step + 1) == original_batch_size + def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> SpecDecodePlan: pre_draft_step = self.pre_draft_step if req_num == 0: self.pre_draft_step = self.max_draft_step - return 0, self.max_draft_step, pre_draft_step - if not self.target_infer_costs.has_data() or not self.draft_infer_costs.has_data(): - # The cost model is only meaningful after both target and draft - # decode costs have been profiled. Block proposers such as DFlash - # do not run through draft_model.forward, and cudagraph may also be - # disabled, so a missing table must not collapse dynamic scheduling to - # draft_step=0. - self.pre_draft_step = self.max_draft_step - return req_num * (pre_draft_step + 1), self.max_draft_step, pre_draft_step - - self._iter += 1 - if self._use_random_mode and self._iter % self._iter_threshold == 0: - min_batch_size = req_num - max_batch_size = req_num * (pre_draft_step + 1) - dynamic_batch_size = self._random.randint(min_batch_size, max_batch_size) + return SpecDecodePlan( + dynamic_batch_size=0, + draft_step=self.max_draft_step, + pre_draft_step=pre_draft_step, + ) - draft_step = self._random.randint(0, self.max_draft_step) - self.pre_draft_step = draft_step - return dynamic_batch_size, draft_step, pre_draft_step + # A request entering its first decode has only the guaranteed target + # row. Existing requests additionally expose pre_draft_step candidates. + # This bounds Verify by the candidate pool that physically exists. + available_batch_size = req_num + proposal_req_num * pre_draft_step + max_batch_size = min(original_batch_size, available_batch_size) + record_progress = proposal_req_num == req_num + + if ( + not self.target_infer_costs.has_data() + or not self.draft_infer_costs.has_data() + or not self.progress_ema_by_config + ): + # Costs come from CUDA Graph capture, while progress requires one + # real verification batch. Keep the available proposal intact until + # both parts of J(N, B, d) have an observation. + self.pre_draft_step = self.max_draft_step + return SpecDecodePlan( + dynamic_batch_size=max_batch_size, + draft_step=self.max_draft_step, + pre_draft_step=pre_draft_step, + record_progress=record_progress, + ) min_batch_size = req_num - max_batch_size = req_num * (pre_draft_step + 1) - dynamic_batch_size_keys = self.target_infer_costs.get_batch_size_keys_between(min_batch_size, max_batch_size) - - cost_ms_list = [ - self._get_cost_ms(req_num=req_num, dynamic_batch_size=dynamic_batch_size, draft_step=pre_draft_step) - for dynamic_batch_size in dynamic_batch_size_keys + batch_sizes = self.target_infer_costs.get_batch_size_keys_between(min_batch_size, max_batch_size) + costs = [ + self._get_cost_ms( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=pre_draft_step, + ) + for dynamic_batch_size in batch_sizes ] - dynamic_batch_size = dynamic_batch_size_keys[np.argmin(cost_ms_list)] - - min_cost_ms = float("inf") - min_cost_ms_draft_step = 0 - for draft_step in range(0, self.max_draft_step + 1): - cost_ms = self._get_cost_ms(req_num=req_num, dynamic_batch_size=dynamic_batch_size, draft_step=draft_step) - if cost_ms < min_cost_ms: - min_cost_ms = cost_ms - min_cost_ms_draft_step = draft_step - - self.pre_draft_step = min_cost_ms_draft_step - return dynamic_batch_size, min_cost_ms_draft_step, pre_draft_step - - def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: - accept_ratio = self._get_dynamic_batch_size_to_accept_ratio( - req_num=req_num, dynamic_batch_size=dynamic_batch_size - ) - total_time = ( - self.target_infer_costs.get(dynamic_batch_size) - + self.draft_infer_costs.get(dynamic_batch_size) * draft_step - ) - token_num = min((dynamic_batch_size * accept_ratio), req_num * (draft_step + 1)) - token_num = max(token_num, req_num) - cost_ms = total_time / token_num - return cost_ms - - def _get_dynamic_batch_size_to_accept_ratio(self, req_num: int, dynamic_batch_size: int): - ema = self._get_req_num_to_dynamic_batch_size_to_accept_ratio( - req_num=req_num, dynamic_batch_size=dynamic_batch_size + dynamic_batch_size = batch_sizes[np.argmin(costs)] + draft_step = self._select_draft_step( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, ) - if ema.get_count() >= 10: - # 当以及通过充分的数据统计以后,直接返回统计的接受率 - return ema.get() - - # 通过单请求的信息进行估计。 - real_step = dynamic_batch_size / req_num - assert real_step >= 1.0 - real_step = real_step - 1.0 - - # 用插值的方式估计不同draft_len 对应的接受率 - left = int(math.floor(real_step)) - right = int(left + 1) - if left == 0: - left_value = 0.0 - else: - left_value = self.draft_len_to_accept_ratio[left - 1].get() - - if right > self.max_draft_step: - right_value = 0.0 - else: - right_value = self.draft_len_to_accept_ratio[right - 1].get() - - accept_ratio = left_value + (right_value - left_value) * (real_step - left) - estimated_accept_ratio = (req_num + (dynamic_batch_size - req_num) * accept_ratio) / dynamic_batch_size - weight = ema.get_count() / 10 - # 通过统计数据和单请求数据进行加权平均,得到最终的接受率 - return estimated_accept_ratio * (1 - weight) + ema.get() * weight - - -class Eagle3DynamicSpecPlanner(DynamicSpecPlanner): - """Joint draft-length and verify-capacity planner for Eagle3. - - ``pre_draft_step`` bounds the proposal that is being verified now, while - ``draft_step`` controls the proposal built for the next target iteration. - Treating those as the same iteration makes ``draft_step == 0`` an - absorbing state. We instead choose the next draft length from a - steady-state search, then choose the current verify capacity within the - proposal width that is actually available now. - - Eagle3 also has an asymmetric draft cost: its first forward commits all K - selected target rows, while every recurrent forward after that processes - one accepted tail per request (batch B). - """ - - _ACCEPT_RATIO_BUCKETS_PER_DRAFT_ROW = 8 - def __init__(self, max_draft_step: int) -> None: - super().__init__(max_draft_step=max_draft_step, use_random_mode=False) - # Eagle uses these values as a full-verify survival curve. Start from - # the first batch mean instead of decaying slowly from an all-accepted - # prior; otherwise a 32-iteration calibration still substantially - # overestimates short draft depths. - self._prefix_survival_decay = float(os.getenv("LIGHTLLM_EAGLE3_PREFIX_SURVIVAL_DECAY", "0.95")) - self.draft_len_to_accept_ratio = [ - _EMAValue(decay=self._prefix_survival_decay, init_value=1.0, enable_decay_warmup=True) - for _ in range(self.max_draft_step) - ] - self._min_static_progress_ratio = float(os.getenv("LIGHTLLM_EAGLE3_MIN_STATIC_PROGRESS_RATIO", "0.85")) - # This is the externally visible verify efficiency: - # accepted target rows / selected target verify rows. It includes the - # guaranteed first row, matching the project acceptance reported by - # the HTTP metrics and benchmark helper. - self._min_project_accept_ratio = float( - os.getenv( - "LIGHTLLM_EAGLE3_MIN_PROJECT_ACCEPT_RATIO", - # Keep a control margin above the externally requested 80% - # acceptance. Stop sequences can discard already-verified - # tail tokens, so HTTP output/verify metrics are about two - # points below the planner's model-accepted/verify feedback - # on GSM8K. - os.getenv("LIGHTLLM_EAGLE3_MIN_DRAFT_ACCEPT_RATIO", "0.860"), - ) - ) - self._full_verify_warmup_steps = max( - 0, - int(os.getenv("LIGHTLLM_EAGLE3_FULL_VERIFY_WARMUP_STEPS", "32")), - ) - self._full_verify_interval = max( - 0, - int(os.getenv("LIGHTLLM_EAGLE3_FULL_VERIFY_INTERVAL", "128")), - ) - self._progress_relax_ratio = float(os.getenv("LIGHTLLM_EAGLE3_PROGRESS_RELAX_RATIO", "1.0")) - self._capacity_accept_ratio_floor = float(os.getenv("LIGHTLLM_EAGLE3_CAPACITY_ACCEPT_RATIO_FLOOR", "0.80")) - self._capacity_feedback_gain = float(os.getenv("LIGHTLLM_EAGLE3_CAPACITY_FEEDBACK_GAIN", "0.10")) - self._capacity_feedback_reference_req_num = max( - 1, - int(os.getenv("LIGHTLLM_EAGLE3_CAPACITY_FEEDBACK_REFERENCE_REQ_NUM", "128")), - ) - self._align_verify_rows_to_graph = os.getenv("LIGHTLLM_EAGLE3_ALIGN_VERIFY_ROWS_TO_GRAPH", "1").lower() in { - "1", - "true", - "yes", - "on", - } - self._max_dynamic_draft_step = max( - 0, - min( - self.max_draft_step, - int(os.getenv("LIGHTLLM_EAGLE3_MAX_DYNAMIC_DRAFT_STEP", str(self.max_draft_step))), - ), - ) - self._full_verify_baseline_decay = float( - os.getenv( - "LIGHTLLM_EAGLE3_BATCH_ACCEPT_EMA_DECAY", - os.getenv("LIGHTLLM_EAGLE3_FULL_VERIFY_BASELINE_DECAY", "0.80"), - ) - ) - assert 0.0 < self._min_static_progress_ratio <= 1.0 - assert 0.0 < self._min_project_accept_ratio <= 1.0 - assert 0.0 < self._progress_relax_ratio <= 1.0 - assert 0.0 < self._capacity_accept_ratio_floor <= 1.0 - assert 0.0 < self._capacity_feedback_gain <= 1.0 - assert 0.0 <= self._prefix_survival_decay < 1.0 - assert 0.0 <= self._full_verify_baseline_decay < 1.0 - - # Exact (B, K) statistics are sparse because the live request batch B - # changes constantly. Pool observations by normalized selected draft - # rows per request so adjacent concurrency levels share evidence. - self._accept_ratio_by_depth_and_width_bucket: Dict[Tuple[int, int], _EMAValue] = {} - self._accept_ratio_by_depth_req_and_batch: Dict[Tuple[int, int, int], _EMAValue] = {} - - # A few full-width verifies provide an unbiased estimate of the static - # Eagle acceptance length. Dynamic top-K observations alone are - # intentionally biased toward high-confidence rows and cannot serve as - # a static baseline. - self._full_verify_tokens_per_req_ema = _EMAValue( - decay=self._full_verify_baseline_decay, - init_value=float(self.max_draft_step + 1), - enable_decay_warmup=True, - ) - self._full_verify_tokens_per_req_value = float(self.max_draft_step + 1) - self._full_verify_accepted_token_sum = 0.0 - self._full_verify_request_count = 0 - self._full_verify_update_count = 0 - self._full_probe_pending = False - - # This feedback comes from real non-full dynamic iterations. It - # provides a conservative capacity floor when a sparse cost-table - # candidate has an over-optimistic expected-token estimate. - self._observed_dynamic_project_accept_ratio = _EMAValue( - decay=0.9, - init_value=self._min_project_accept_ratio, - enable_decay_warmup=True, - ) - self._observed_dynamic_tokens_per_req = _EMAValue( - decay=0.8, - init_value=1.0, - enable_decay_warmup=True, + self.pre_draft_step = draft_step + return SpecDecodePlan( + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + pre_draft_step=pre_draft_step, + record_progress=record_progress, ) - # Acceptance is aggregated with request weights. An iteration with a - # single long-tail request must not have the same influence as an - # iteration with 256 live requests. - self._observed_dynamic_accepted_token_sum = 0.0 - self._observed_dynamic_verify_row_sum = 0.0 - self._observed_dynamic_request_count = 0 - - # Closed-loop verify capacity. The initial value comes from the - # unbiased full-width baseline. Real dynamic iterations then move it - # toward the largest width that still satisfies project acceptance. - self._target_verify_rows_per_req_value: Optional[float] = None - - self._plan_count = 0 - - def update_verified_prefix_stats(self, verify_len: int, accept_len: int) -> None: - """Record the survival probability of each Eagle draft position. - - The generic planner records ``(accept_len - 1) / depth``. That value - is neither a conditional probability nor a survival probability and - can increase with depth. Eagle's expected-token model needs the - probability that a token at each depth is actually reached. - """ - - max_draft_len = min(max(0, verify_len - 1), self.max_draft_step) - for draft_len in range(1, max_draft_len + 1): - self.update_draft_len_to_accept_ratio( - draft_len=draft_len, - accept_ratio=1.0 if accept_len > draft_len else 0.0, - ) - def update_verified_batch_prefix_stats( - self, - verify_and_accept_lengths: List[Tuple[int, int]], - ) -> None: - """Update each depth once with a request-weighted batch mean.""" - - for draft_len in range(1, self.max_draft_step + 1): - eligible_accept_lengths = [ - accept_len for verify_len, accept_len in verify_and_accept_lengths if verify_len > draft_len - ] - if not eligible_accept_lengths: - continue - survival_ratio = sum(accept_len > draft_len for accept_len in eligible_accept_lengths) / len( - eligible_accept_lengths - ) - self.update_draft_len_to_accept_ratio( - draft_len=draft_len, - accept_ratio=survival_ratio, - ) + def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: + cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs + cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - def update_req_num_to_dynamic_batch_size_to_accept_ratio( + def update_feedback( self, + plan: SpecDecodePlan, req_num: int, - dynamic_batch_size: int, - accept_ratio: float, - pre_draft_step: int = None, - ) -> None: - depth = self.max_draft_step if pre_draft_step is None else int(pre_draft_step) - exact_key = (depth, int(req_num), int(dynamic_batch_size)) - if exact_key not in self._accept_ratio_by_depth_req_and_batch: - self._accept_ratio_by_depth_req_and_batch[exact_key] = _EMAValue( - decay=0.9, - init_value=1.0, - enable_decay_warmup=True, - ) - self._accept_ratio_by_depth_req_and_batch[exact_key].update(accept_ratio) - self._get_width_bucket_accept_ratio( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - pre_draft_step=pre_draft_step, - ).update(accept_ratio) - - def update_full_verify_tokens_per_req(self, tokens_per_req: float, req_num: int = 1) -> None: - tokens_per_req = max(1.0, min(float(self.max_draft_step + 1), float(tokens_per_req))) - req_num = max(1, int(req_num)) - self._full_verify_accepted_token_sum += tokens_per_req * req_num - self._full_verify_request_count += req_num - # A lifetime cumulative average cannot follow a non-stationary trace: - # a preceding predictable segment would keep the progress floor high - # long after acceptance collapses. Track the live full-width baseline - # with an iteration-level EMA. Request counts remain cumulative only - # for diagnostics. - self._full_verify_tokens_per_req_ema.update(tokens_per_req) - self._full_verify_tokens_per_req_value = self._full_verify_tokens_per_req_ema.get() - self._full_verify_update_count += 1 - - def update_observed_iteration_stats( - self, - tokens_per_req: float, - verify_rows_per_req: float, - is_full_verify: bool, - req_num: int = 1, + accept_lengths, + schedule_scores=None, ) -> None: - if is_full_verify or verify_rows_per_req <= 1.0: + # The progress EMA records one complete-batch sample for a single + # (N, B, d) configuration. record_progress is false if any request is + # on its first decode and therefore has no preceding proposal; its + # structural accept_len=1 would otherwise bias that configuration's + # progress downward. Skip the mixed batch instead of partially + # updating the batch-level statistic with different N/B semantics. + if not plan.record_progress: return - req_num = max(1, int(req_num)) - project_accept_ratio = max(0.0, min(1.0, tokens_per_req / verify_rows_per_req)) - self._observed_dynamic_project_accept_ratio.update(project_accept_ratio) - self._observed_dynamic_tokens_per_req.update(tokens_per_req) - self._observed_dynamic_accepted_token_sum += tokens_per_req * req_num - self._observed_dynamic_verify_row_sum += verify_rows_per_req * req_num - self._observed_dynamic_request_count += req_num - - self._update_target_verify_rows_per_req( - tokens_per_req=tokens_per_req, - verify_rows_per_req=verify_rows_per_req, + self.update_verified_batch( + accept_lengths=accept_lengths, req_num=req_num, + dynamic_batch_size=plan.dynamic_batch_size, + verified_draft_step=plan.pre_draft_step, ) - def _update_target_verify_rows_per_req( - self, - tokens_per_req: float, - verify_rows_per_req: float, - req_num: int, - ) -> None: - current_target = self._get_controlled_verify_rows_per_req() - - # Holding the accepted-token count locally constant, this is the - # verify width that lands exactly on the project-acceptance target. - sample_target = tokens_per_req / self._min_project_accept_ratio - - # If progress is below its static-relative floor, acceptance and - # progress constraints conflict. Widen enough to recover progress; - # subsequent observations will pull the controller back once the - # additional rows stop paying for themselves. - min_expected_tokens = self._get_min_expected_tokens_per_req() - if self._get_observed_tokens_per_req() < min_expected_tokens: - sample_target = max( - sample_target, - verify_rows_per_req * min_expected_tokens / max(tokens_per_req, 1e-6), - ) - - sample_target = max(1.0, min(float(self.max_draft_step + 1), sample_target)) - # Bound a single observation. This is especially important while a - # batch drains and only a handful of unusually hard requests remain. - sample_target = max(current_target - 0.5, min(current_target + 0.5, sample_target)) - request_weight = req_num / self._capacity_feedback_reference_req_num - alpha = 1.0 - (1.0 - self._capacity_feedback_gain) ** request_weight - self._target_verify_rows_per_req_value = current_target * (1.0 - alpha) + sample_target * alpha - - def _get_observed_project_accept_ratio(self) -> float: - if self._observed_dynamic_verify_row_sum <= 0.0: - return self._min_project_accept_ratio - return self._observed_dynamic_accepted_token_sum / self._observed_dynamic_verify_row_sum - - def _get_observed_tokens_per_req(self) -> float: - if self._observed_dynamic_request_count <= 0: - return self._observed_dynamic_tokens_per_req.get() - return self._observed_dynamic_accepted_token_sum / self._observed_dynamic_request_count - - def _get_dynamic_batch_size_to_accept_ratio( + def update_verified_batch( self, + accept_lengths, req_num: int, dynamic_batch_size: int, - pre_draft_step: int = None, - ): - depth = self.max_draft_step if pre_draft_step is None else int(pre_draft_step) - exact_key = (depth, int(req_num), int(dynamic_batch_size)) - exact_ema = self._accept_ratio_by_depth_req_and_batch.get(exact_key) - base_estimate = super()._get_dynamic_batch_size_to_accept_ratio( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - ) - if exact_ema is not None and exact_ema.get_count() >= 10: - return exact_ema.get() + verified_draft_step: int, + ) -> None: + """Record one batch-level progress sample for the verified configuration.""" - width_ema = self._get_width_bucket_accept_ratio( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - pre_draft_step=pre_draft_step, - ) - width_weight = min(1.0, width_ema.get_count() / 10.0) - return base_estimate * (1.0 - width_weight) + width_ema.get() * width_weight + accept_lengths = np.asarray(accept_lengths) + if accept_lengths.size == 0: + return - def _get_width_bucket_accept_ratio( - self, - req_num: int, - dynamic_batch_size: int, - pre_draft_step: int = None, - ) -> "_EMAValue": - selected_draft_rows_per_req = max(0.0, dynamic_batch_size / req_num - 1.0) - width_bucket = int(round(selected_draft_rows_per_req * self._ACCEPT_RATIO_BUCKETS_PER_DRAFT_ROW)) - depth = self.max_draft_step if pre_draft_step is None else int(pre_draft_step) - bucket = (depth, width_bucket) - if bucket not in self._accept_ratio_by_depth_and_width_bucket: - self._accept_ratio_by_depth_and_width_bucket[bucket] = _EMAValue( + config = (int(req_num), int(dynamic_batch_size), int(verified_draft_step)) + progress = float(accept_lengths.sum()) / dynamic_batch_size + if config not in self.progress_ema_by_config: + init_progress = progress if not self.progress_ema_by_config else self._estimate_nearby_progress(*config) + self.progress_ema_by_config[config] = _EMAValue( decay=0.9, - init_value=1.0, - enable_decay_warmup=True, + init_value=init_progress, ) - return self._accept_ratio_by_depth_and_width_bucket[bucket] - - def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: - assert req_num * (self.max_draft_step + 1) == original_batch_size - pre_draft_step = self.pre_draft_step + self.progress_ema_by_config[config].update(progress) - if req_num == 0: - self.pre_draft_step = self.max_draft_step - return 0, self.max_draft_step, pre_draft_step - - max_batch_size = req_num * (pre_draft_step + 1) - if not self.target_infer_costs.has_data() or not self.draft_infer_costs.has_data(): - self.pre_draft_step = self.max_draft_step - return max_batch_size, self.max_draft_step, pre_draft_step - - if self._should_schedule_full_probe(): - self._full_probe_pending = True - - force_full_verify = self._should_force_full_verify(pre_draft_step=pre_draft_step) - if force_full_verify: - # Keep drafting full width during initial calibration. A periodic - # probe may return to the normal dynamic draft length immediately - # after its one full target verify. - draft_step = ( - self.max_draft_step - if self._full_verify_update_count < self._full_verify_warmup_steps - else self._select_next_draft_step(req_num=req_num) - ) - dynamic_batch_size = max_batch_size - self.pre_draft_step = draft_step - self._record_plan() - return dynamic_batch_size, draft_step, pre_draft_step - - # Once the full-width baseline is calibrated, choose draft depth by - # the target+draft cost model, but control verify width with real - # project-acceptance feedback. Sparse (B, K) cost/acceptance buckets - # are too optimistic for unseen widths and previously widened GSM8K - # from about 4.5 to 6 rows/request before converging. - if self._full_verify_request_count > 0: - draft_step = self._select_next_draft_step(req_num=req_num) - if self._full_probe_pending: - draft_step = self.max_draft_step - target_verify_rows_per_req = self._get_controlled_verify_rows_per_req() - dynamic_batch_size = int(math.ceil(req_num * target_verify_rows_per_req)) - if self._align_verify_rows_to_graph: - # Replay already pads an arbitrary target batch to the next - # captured shape. Turn those paid-for padding rows into real - # candidates so they can improve accepted progress for the - # same target-model graph cost. - graph_batch_size = self.target_infer_costs.get_ceil_batch_size( - dynamic_batch_size, - max_batch_size=max_batch_size, - ) - if graph_batch_size is not None: - dynamic_batch_size = graph_batch_size - dynamic_batch_size = min(max(dynamic_batch_size, req_num), max_batch_size) - self.pre_draft_step = draft_step - self._record_plan() - return dynamic_batch_size, draft_step, pre_draft_step - - # The action selected here builds the proposal consumed by the next - # target iteration. Search it independently of the current proposal - # width so a zero-step iteration can recover on its own. - draft_step = self._select_next_draft_step(req_num=req_num) - if self._full_probe_pending: - draft_step = self.max_draft_step - - dynamic_batch_size_keys = self._get_candidate_batch_sizes( - req_num=req_num, - max_batch_size=max_batch_size, - ) - candidates = [ - self._get_eagle3_candidate( + def _select_draft_step(self, req_num: int, dynamic_batch_size: int) -> int: + best_cost_ms = float("inf") + best_draft_step = self.draft_steps[0] + for draft_step in self.draft_steps: + cost_ms = self._get_cost_ms( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - pre_draft_step=pre_draft_step, draft_step=draft_step, ) - for dynamic_batch_size in dynamic_batch_size_keys - ] - best_candidate = self._select_best_candidate(candidates, req_num=req_num) - dynamic_batch_size = best_candidate[1] - self.pre_draft_step = draft_step - self._record_plan() - return dynamic_batch_size, draft_step, pre_draft_step - - def _get_target_verify_rows_per_req(self) -> float: - min_expected_tokens = self._get_min_expected_tokens_per_req() - if min_expected_tokens <= 1.0: - return 1.0 - return min( - float(self.max_draft_step + 1), - min_expected_tokens / self._min_project_accept_ratio, - ) - - def _get_controlled_verify_rows_per_req(self) -> float: - if self._target_verify_rows_per_req_value is None: - return self._get_target_verify_rows_per_req() - return max( - 1.0, - min(float(self.max_draft_step + 1), self._target_verify_rows_per_req_value), - ) + if cost_ms < best_cost_ms: + best_cost_ms = cost_ms + best_draft_step = draft_step + return best_draft_step - def _get_candidate_batch_sizes(self, req_num: int, max_batch_size: int) -> List[int]: - candidates = set( - self.target_infer_costs.get_batch_size_keys_between( - req_num, - max_batch_size, - ) - ) - candidates.add(req_num) - candidates.add(max_batch_size) - exact_target = int(math.ceil(req_num * self._get_target_verify_rows_per_req())) - if req_num <= exact_target <= max_batch_size: - candidates.add(exact_target) - - # Add exact points along the feasible progress/acceptance frontier. - # CUDA graph timing keys are deliberately coarse; without these - # points the cost search can only choose between two widely separated - # capacities and often leaves useful acceptance headroom unused. - if self._full_verify_request_count > 0: - for progress_ratio in np.linspace(self._min_static_progress_ratio, 1.0, num=7): - expected_tokens_per_req = self._full_verify_tokens_per_req_value * float(progress_ratio) - verify_rows_per_req = expected_tokens_per_req / self._min_project_accept_ratio - frontier_batch_size = int(math.ceil(req_num * verify_rows_per_req)) - if req_num <= frontier_batch_size <= max_batch_size: - candidates.add(frontier_batch_size) - return sorted(candidates) - - def _should_schedule_full_probe(self) -> bool: - return ( - self._full_verify_interval > 0 - and self._full_verify_update_count >= self._full_verify_warmup_steps - and self._plan_count > 0 - and self._plan_count % self._full_verify_interval == 0 - ) - - def _should_force_full_verify(self, pre_draft_step: int) -> bool: - if pre_draft_step != self.max_draft_step: - return False - if self._full_verify_update_count < self._full_verify_warmup_steps: - return True - if self._full_probe_pending: - self._full_probe_pending = False - return True - return False - - def _record_plan(self) -> None: - self._plan_count += 1 - - def _select_next_draft_step(self, req_num: int) -> int: - """Choose a recoverable long-run Eagle3 draft length. - - For each possible length, jointly search the verify capacities that - would be legal if that length were used repeatedly. This is a - one-state steady approximation to the next iteration and avoids - crediting newly drafted tokens to the current target forward. - """ + def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: + """Estimate milliseconds per committed token for one configuration.""" - candidates = [] - min_expected_tokens = self._get_min_expected_tokens_per_req() - for draft_step in range(self._max_dynamic_draft_step + 1): - # Even a full verify at this proposal depth must be capable of - # meeting the static-relative progress floor. This prevents an - # optimistic sparse (B, K) bucket from selecting draft_step=3 on - # GSM8K and paying for many extra target iterations. - max_depth_tokens_per_req = 1.0 + sum( - self.draft_len_to_accept_ratio[draft_index].get() for draft_index in range(draft_step) - ) - if max_depth_tokens_per_req + 1e-6 < min_expected_tokens: - continue - max_batch_size = req_num * (draft_step + 1) - dynamic_batch_size_keys = self._get_candidate_batch_sizes( - req_num=req_num, - max_batch_size=max_batch_size, - ) - for dynamic_batch_size in dynamic_batch_size_keys: - candidates.append( - self._get_eagle3_candidate( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - pre_draft_step=draft_step, - draft_step=draft_step, - ) - ) - if not candidates: - return self._max_dynamic_draft_step - return self._select_best_candidate(candidates, req_num=req_num)[4] - - def _get_eagle3_candidate( - self, - req_num: int, - dynamic_batch_size: int, - pre_draft_step: int, - draft_step: int, - ) -> Tuple[float, int, float, float, int]: - expected_token_num = self._estimate_expected_token_num( + accept_ratio = self._estimate_progress( req_num=req_num, dynamic_batch_size=dynamic_batch_size, - pre_draft_step=pre_draft_step, + draft_step=draft_step, ) - cost_ms = self._get_eagle3_cost_ms( + total_time = self.target_infer_costs.get(dynamic_batch_size) + self.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=self.draft_infer_costs, req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - pre_draft_step=pre_draft_step, + verify_batch_size=dynamic_batch_size, draft_step=draft_step, - expected_token_num=expected_token_num, - ) - expected_tokens_per_req = expected_token_num / req_num - project_accept_ratio = max(0.0, min(1.0, expected_token_num / dynamic_batch_size)) - return ( - cost_ms, - dynamic_batch_size, - expected_tokens_per_req, - project_accept_ratio, - draft_step, ) + token_num = min((dynamic_batch_size * accept_ratio), req_num * (draft_step + 1)) + token_num = max(token_num, req_num) + cost_ms = total_time / token_num + return cost_ms - def _select_best_candidate( - self, - candidates: List[Tuple[float, int, float, float, int]], - req_num: int, - ) -> Tuple[float, int, float, float, int]: - assert candidates - min_expected_tokens = self._get_min_expected_tokens_per_req() - min_verify_rows = self._get_min_verify_rows_per_req() - progress_candidates = [ - candidate - for candidate in candidates - if candidate[2] >= min_expected_tokens and candidate[1] / req_num >= min_verify_rows - ] - efficient_candidates = [ - candidate for candidate in progress_candidates if candidate[3] >= self._min_project_accept_ratio - ] - if efficient_candidates: - return min(efficient_candidates, key=lambda candidate: candidate[0]) - - # If noisy online estimates leave no candidate satisfying both hard - # constraints, minimize their worst relative violation. This avoids - # collapsing to the narrow acceptance-only plan or jumping to the - # wide progress-only plan while the exact boundary bucket converges. - drafted_candidates = [candidate for candidate in candidates if candidate[1] > req_num] - if drafted_candidates: - best_constraint_score = max( - min( - candidate[2] / max(min_expected_tokens, 1e-6), - candidate[3] / max(self._min_project_accept_ratio, 1e-6), - ) - for candidate in drafted_candidates - ) - balanced_candidates = [ - candidate - for candidate in drafted_candidates - if min( - candidate[2] / max(min_expected_tokens, 1e-6), - candidate[3] / max(self._min_project_accept_ratio, 1e-6), - ) - >= best_constraint_score - 1e-6 - ] - return min(balanced_candidates, key=lambda candidate: candidate[0]) - - # The current proposal may be too short to meet the floor after a - # transition or drafting may be disabled. Make maximal forward - # progress instead of falling back to a deceptively cheap iteration. - max_progress = max(candidate[2] for candidate in candidates) - max_progress_candidates = [candidate for candidate in candidates if candidate[2] >= max_progress - 1e-6] - return min(max_progress_candidates, key=lambda candidate: candidate[0]) - - def _get_min_expected_tokens_per_req(self) -> float: - if self._full_verify_request_count == 0: - return 1.0 - target_tokens = max( - 1.0, - self._full_verify_tokens_per_req_value * self._min_static_progress_ratio, - ) - if self._observed_dynamic_request_count > 0 and self._get_observed_tokens_per_req() >= target_tokens: - return max(1.0, target_tokens * self._progress_relax_ratio) - return target_tokens + def _estimate_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: + config = (int(req_num), int(dynamic_batch_size), int(draft_step)) + ema = self.progress_ema_by_config.get(config) + return ema.get() if ema is not None else self._estimate_nearby_progress(*config) - def _get_min_verify_rows_per_req(self) -> float: - min_expected_tokens = self._get_min_expected_tokens_per_req() - if min_expected_tokens <= 1.0: + def _estimate_nearby_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: + if not self.progress_ema_by_config: return 1.0 - observed_project_accept_ratio = max( - self._capacity_accept_ratio_floor, - self._get_observed_project_accept_ratio(), - ) - return min_expected_tokens / observed_project_accept_ratio - def _estimate_expected_token_num(self, req_num: int, dynamic_batch_size: int, pre_draft_step: int) -> float: - accept_ratio = self._get_dynamic_batch_size_to_accept_ratio( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - pre_draft_step=pre_draft_step, - ) - expected_token_num = min( - dynamic_batch_size * accept_ratio, - req_num * (pre_draft_step + 1), + rows_per_req = dynamic_batch_size / req_num + + def distance(config: Tuple[int, int, int]): + observed_req_num, observed_batch_size, observed_draft_step = config + width_distance = (observed_batch_size / observed_req_num - rows_per_req) / (self.max_draft_step + 1) + draft_distance = (observed_draft_step - draft_step) / max(self.max_draft_step, 1) + request_distance = (observed_req_num - req_num) / max(observed_req_num, req_num) + return width_distance ** 2 + draft_distance ** 2 + request_distance ** 2 + + nearest_config = min(self.progress_ema_by_config, key=distance) + observed_req_num, observed_batch_size, _ = nearest_config + observed_rows_per_req = observed_batch_size / observed_req_num + observed_progress_per_req = observed_rows_per_req * self.progress_ema_by_config[nearest_config].get() + + # Fill keeps the highest-survival rows when B shrinks. Transfer the + # neighbor's committed progress per request, then normalize it for the + # target B. Copying rho directly would predict that useful progress + # falls in proportion to B and make B=N an artificial absorbing state. + estimated_progress_per_req = min( + max(observed_progress_per_req, 1.0), + rows_per_req, + draft_step + 1, ) - return max(float(req_num), expected_token_num) + return estimated_progress_per_req / rows_per_req - def _get_eagle3_cost_ms( - self, - req_num: int, - dynamic_batch_size: int, - pre_draft_step: int, - draft_step: int, - expected_token_num: float = None, - ) -> float: - if expected_token_num is None: - expected_token_num = self._estimate_expected_token_num( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - pre_draft_step=pre_draft_step, - ) - # Eagle3 commit runs on all selected verify rows. Every recurrent - # proposal step after that runs on one accepted tail per request. - draft_cost_ms = 0.0 - if draft_step > 0: - draft_cost_ms = self.draft_infer_costs.get(dynamic_batch_size) - if draft_step > 1: - draft_cost_ms += self.draft_infer_costs.get(req_num) * (draft_step - 1) +class DSparkPlanner: + """DSpark's confidence-based verify-capacity planner. - total_time_ms = self.target_infer_costs.get(dynamic_batch_size) + draft_cost_ms - return total_time_ms / expected_token_num + DSpark always drafts a complete block. Confidence from one proposal selects + the target verify capacity two iterations later. + """ + def __init__(self, max_draft_step: int, draft_cost_provider: "BaseSpecProposer") -> None: + self.max_draft_step = int(max_draft_step) + self.draft_cost_provider = draft_cost_provider + self.target_infer_costs = _InferCostMsTable() + self.draft_infer_costs = _InferCostMsTable() + self._pending_verify_batch_sizes = deque(maxlen=2) -class DSparkDynamicSpecPlanner(DynamicSpecPlanner): - """DSpark confidence-scheduled verify-capacity planner.""" + def plan(self, req_num: int, original_batch_size: int) -> SpecDecodePlan: + if req_num == 0: + return SpecDecodePlan( + dynamic_batch_size=0, + draft_step=self.max_draft_step, + pre_draft_step=self.max_draft_step, + ) - def __init__(self, max_draft_step: int) -> None: - super().__init__( - max_draft_step=max_draft_step, - use_random_mode=False, + full_batch_size = original_batch_size + dynamic_batch_size = full_batch_size + if self.target_infer_costs.has_data() and self.draft_infer_costs.has_data(): + delayed_batch_size = self._pop_delayed_batch_size( + req_num=req_num, + max_batch_size=full_batch_size, + ) + if delayed_batch_size is not None: + dynamic_batch_size = delayed_batch_size + + return SpecDecodePlan( + dynamic_batch_size=dynamic_batch_size, + draft_step=self.max_draft_step, + pre_draft_step=self.max_draft_step, ) - self.draft_len_to_accept_ratio = [ - _EMAValue(decay=0.95, init_value=1.0, enable_decay_warmup=True) for _ in range(self.max_draft_step) - ] - self._predicted_dynamic_batch_sizes = deque(maxlen=2) - def update_verified_prefix_stats(self, verify_len: int, accept_len: int) -> None: - if verify_len - 1 <= 0: - return - max_draft_len = min(verify_len - 1, self.max_draft_step) - for draft_len in range(1, max_draft_len + 1): - self.update_draft_len_to_accept_ratio( - draft_len=draft_len, - accept_ratio=1.0 if accept_len > draft_len else 0.0, + def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: + cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs + cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + + def update_feedback( + self, + plan: SpecDecodePlan, + req_num: int, + accept_lengths, + schedule_scores=None, + ) -> None: + if schedule_scores is not None: + self.update_confidence_probs( + confidence_probs=schedule_scores, + req_num=req_num, ) - def update_predicted_schedule_probs(self, schedule_probs, req_num: int) -> None: + def update_confidence_probs(self, confidence_probs, req_num: int) -> None: """Record a confidence-derived future capacity estimate. The current decode step still routes rows by the current probabilities - stored in req_to_next_token_probs. This queue is only used to choose + stored in req_to_next_token_scores. This queue is only used to choose the future capacity K after a two-step delay, matching the asynchronous scheduler constraint described by DSpark. """ if req_num <= 0: return - if not self.target_infer_costs.has_data(): + if not self.target_infer_costs.has_data() or not self.draft_infer_costs.has_data(): return - probs = self._to_numpy(schedule_probs) - if probs is None or probs.ndim != 2 or probs.shape[1] <= 1: + probs = np.asarray(confidence_probs, dtype=np.float64) + if probs.ndim != 2 or probs.shape[1] == 0: return - draft_probs = probs[:, 1 : self.max_draft_step + 1] - if draft_probs.size == 0: + draft_confidence_probs = probs[:, : self.max_draft_step] + if draft_confidence_probs.size == 0: return - valid_rows = np.any(draft_probs > 0.0, axis=1) + # Confidence is scattered onto one accepted-tail row per request; + # unused verify rows remain zero. + valid_rows = np.any(draft_confidence_probs > 0.0, axis=1) if not np.any(valid_rows): return - conditional_probs = np.clip(draft_probs[valid_rows], 0.01, 0.99) + conditional_probs = np.clip(draft_confidence_probs[valid_rows], 0.01, 0.99) survival_scores = np.cumprod(conditional_probs, axis=1) dynamic_batch_size = self._select_dynamic_batch_size_from_survival_scores( req_num=int(req_num), survival_scores=survival_scores, ) - self._predicted_dynamic_batch_sizes.append(dynamic_batch_size) - - def get_dynamic_batch_size(self, req_num: int, original_batch_size: int) -> Tuple[int, int, int]: - assert req_num * (self.max_draft_step + 1) == original_batch_size - pre_draft_step = self.pre_draft_step - self.pre_draft_step = self.max_draft_step - if req_num == 0: - return 0, self.max_draft_step, pre_draft_step - - max_batch_size = req_num * (pre_draft_step + 1) - if not self.target_infer_costs.has_data(): - return max_batch_size, self.max_draft_step, pre_draft_step - - historical_batch_size = self._pop_historical_dynamic_batch_size( - req_num=req_num, - max_batch_size=max_batch_size, - ) - if historical_batch_size is not None: - return historical_batch_size, self.max_draft_step, pre_draft_step - if len(self._predicted_dynamic_batch_sizes) > 0: - # A confidence estimate is available but has not satisfied the - # two-step async delay yet. Keep capacity conservative instead of - # leaking a same-step EMA fallback into DSpark scheduling. - return req_num, self.max_draft_step, pre_draft_step - - candidate_batch_sizes = set(self.target_infer_costs.get_batch_size_keys_between(req_num, max_batch_size)) - candidate_batch_sizes.add(req_num) - candidate_batch_sizes.add(max_batch_size) - survival_prefix = self._estimate_survival_prefix(pre_draft_step) - - best_batch_size = req_num - best_throughput = -float("inf") - for candidate_batch_size in sorted(candidate_batch_sizes): - dynamic_batch_size = min(max(int(candidate_batch_size), req_num), max_batch_size) - expected_tokens = self._estimate_expected_tokens( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - survival_prefix=survival_prefix, - ) - verify_ms = max(self.target_infer_costs.get(dynamic_batch_size), 1e-6) - throughput = expected_tokens / verify_ms - if throughput > best_throughput: - best_throughput = throughput - best_batch_size = dynamic_batch_size + self._pending_verify_batch_sizes.append(dynamic_batch_size) - return best_batch_size, self.max_draft_step, pre_draft_step - - def _pop_historical_dynamic_batch_size(self, req_num: int, max_batch_size: int) -> Optional[int]: - if len(self._predicted_dynamic_batch_sizes) < 2: + def _pop_delayed_batch_size(self, req_num: int, max_batch_size: int) -> Optional[int]: + if len(self._pending_verify_batch_sizes) < 2: return None - predicted_batch_size = int(self._predicted_dynamic_batch_sizes.popleft()) + predicted_batch_size = int(self._pending_verify_batch_sizes.popleft()) return min(max(predicted_batch_size, int(req_num)), int(max_batch_size)) def _select_dynamic_batch_size_from_survival_scores( @@ -1006,8 +405,13 @@ def _select_dynamic_batch_size_from_survival_scores( for dynamic_batch_size in sorted(candidate_batch_sizes): selected_draft_count = dynamic_batch_size - int(req_num) expected_tokens = float(req_num) + float(expected_accepts_by_count[selected_draft_count]) - verify_ms = max(self.target_infer_costs.get(dynamic_batch_size), 1e-6) - throughput = expected_tokens / verify_ms + round_ms = self.target_infer_costs.get(dynamic_batch_size) + self.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=self.draft_infer_costs, + req_num=req_num, + verify_batch_size=dynamic_batch_size, + draft_step=self.max_draft_step, + ) + throughput = expected_tokens / max(round_ms, 1e-6) if throughput > best_throughput: best_throughput = throughput best_batch_size = dynamic_batch_size @@ -1017,74 +421,17 @@ def _select_dynamic_batch_size_from_survival_scores( def _topk_prefix_sums(values: np.ndarray, counts: List[int]) -> Dict[int, float]: """Return sum(top-k(values)) only for the requested k values.""" - if len(counts) == 0: + if not counts: return {} flat_values = np.asarray(values, dtype=np.float64).reshape(-1) value_count = int(flat_values.shape[0]) - ans: Dict[int, float] = {} - normalized_counts = sorted({min(max(int(count), 0), value_count) for count in counts}) if value_count == 0: return {count: 0.0 for count in normalized_counts} - partial_counts = [count for count in normalized_counts if 0 < count < value_count] - if len(partial_counts) > 0 and partial_counts[-1] > value_count // 2: - sorted_values = np.sort(flat_values)[::-1] - prefix_sums = np.cumsum(sorted_values, dtype=np.float64) - for count in normalized_counts: - ans[count] = 0.0 if count == 0 else float(prefix_sums[count - 1]) - return ans - - if normalized_counts[0] == 0: - ans[0] = 0.0 - if normalized_counts[-1] == value_count: - ans[value_count] = float(np.sum(flat_values, dtype=np.float64)) - if len(partial_counts) == 0: - return ans - - max_partial_count = partial_counts[-1] - top_values = np.partition(flat_values, value_count - max_partial_count)[value_count - max_partial_count :] - top_values.sort() - top_values = top_values[::-1] - prefix_sums = np.cumsum(top_values, dtype=np.float64) - for count in partial_counts: - ans[count] = float(prefix_sums[count - 1]) - return ans - - def _estimate_survival_prefix(self, pre_draft_step: int) -> List[float]: - survival_prefix = [1.0] - previous = 1.0 - for draft_index in range(pre_draft_step): - survival = float(self.draft_len_to_accept_ratio[draft_index].get()) - survival = max(0.0, min(previous, survival)) - survival_prefix.append(survival) - previous = survival - return survival_prefix - - @staticmethod - def _estimate_expected_tokens( - req_num: int, - dynamic_batch_size: int, - survival_prefix: List[float], - ) -> float: - expected_tokens = 0.0 - remaining = int(dynamic_batch_size) - for survival in survival_prefix: - take = min(req_num, remaining) - if take <= 0: - break - expected_tokens += take * float(survival) - remaining -= take - return max(float(req_num), expected_tokens) - - @staticmethod - def _to_numpy(value): - if value is None: - return None - if hasattr(value, "detach"): - value = value.detach().cpu().numpy() - return np.asarray(value, dtype=np.float64) + prefix_sums = np.concatenate(([0.0], np.cumsum(np.sort(flat_values)[::-1], dtype=np.float64))) + return {count: float(prefix_sums[count]) for count in normalized_counts} class _InferCostMsTable: @@ -1092,14 +439,12 @@ def __init__(self) -> None: self.infer_cost_ms_table = SortedDict() def update(self, batch_size: int, infer_cost_ms: float) -> None: - assert batch_size > 0 self.infer_cost_ms_table[int(batch_size)] = float(infer_cost_ms) def has_data(self) -> bool: return len(self.infer_cost_ms_table) > 0 def get(self, batch_size: int) -> float: - assert batch_size > 0 batch_size = int(batch_size) if len(self.infer_cost_ms_table) == 0: @@ -1117,54 +462,31 @@ def get(self, batch_size: int) -> float: index = self.infer_cost_ms_table.bisect_left(batch_size) return self.infer_cost_ms_table.peekitem(index)[1] + def estimate(self, batch_size: int) -> float: + """Estimate an uncaptured batch without applying the graph-miss penalty.""" + + batch_size = int(batch_size) + max_batch_size, max_cost_ms = self.infer_cost_ms_table.peekitem(-1) + if batch_size <= max_batch_size: + return self.get(batch_size) + return max_cost_ms * batch_size / max_batch_size + def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> List[int]: - assert batch_size1 > 0 and batch_size2 > 0 start = min(int(batch_size1), int(batch_size2)) end = max(int(batch_size1), int(batch_size2)) batch_sizes = list(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) return batch_sizes or [end] - def get_ceil_batch_size(self, batch_size: int, max_batch_size: int) -> Optional[int]: - """Return the next recorded graph shape without inventing a key. - - The cost table can be sparse or empty when CUDA graph is disabled, so - callers must be able to distinguish "no captured shape" from the - requested upper bound. - """ - - if len(self.infer_cost_ms_table) == 0: - return None - index = self.infer_cost_ms_table.bisect_left(int(batch_size)) - if index >= len(self.infer_cost_ms_table): - return None - candidate = int(self.infer_cost_ms_table.peekitem(index)[0]) - if candidate > int(max_batch_size): - return None - return candidate - class _EMAValue: - def __init__(self, decay: float, init_value: float, enable_decay_warmup: bool = True) -> None: - # decay=0 is the explicit latest-observation ablation used to compare - # EMA-smoothed online estimates against an otherwise identical - # controller. Production defaults remain strictly between zero/one. - assert 0.0 <= decay < 1.0 - self.enable_decay_warmup = enable_decay_warmup + def __init__(self, decay: float, init_value: float) -> None: self.decay = decay - - if self.enable_decay_warmup: - self.current_decay = 0.0 - else: - self.current_decay = self.decay - self.value = init_value self.update_count = 0 def update(self, new_value: float): self.update_count += 1 - self.value = self.current_decay * self.value + (1.0 - self.current_decay) * new_value - # 更新 current_decay 的值,使得 current_decay 逐渐逼近 decay 的值 - self.current_decay = min(self.decay, (self.decay + self.current_decay) / 2.0 + 0.001) + self.value = self.decay * self.value + (1.0 - self.decay) * new_value def get(self) -> float: return self.value @@ -1174,9 +496,8 @@ def get_count(self) -> int: __all__ = [ - "DSparkDynamicSpecPlanner", - "DynamicSpecPlanner", - "Eagle3DynamicSpecPlanner", + "DSparkPlanner", "FixedSpecPlanner", + "LightSpecPlanner", "SpecDecodePlan", ] diff --git a/lightllm/server/router/model_infer/speculative/proposers/__init__.py b/lightllm/server/router/model_infer/speculative/proposers/__init__.py index 85a38a81ca..10f3e31b47 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/__init__.py +++ b/lightllm/server/router/model_infer/speculative/proposers/__init__.py @@ -4,28 +4,27 @@ from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer -def build_spec_proposer(engine) -> "BaseSpecProposer": - spec_mode = engine.spec_mode +def build_spec_proposer(*, spec_mode: str, backend, enable_dynamic_spec: bool) -> "BaseSpecProposer": if spec_mode == "dspark": from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer - return DSparkProposer(engine=engine) + return DSparkProposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) if spec_mode == "dflash": from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer - return DFlashProposer(engine=engine) + return DFlashProposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) if spec_mode == "eagle3": from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer - return Eagle3Proposer(engine=engine) + return Eagle3Proposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) if spec_mode in ("eagle_with_att", "eagle_no_att"): from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer - return EagleMTPProposer(engine=engine) + return EagleMTPProposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) if spec_mode in ("vanilla_with_att", "vanilla_no_att"): from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer - return VanillaMTPProposer(engine=engine) + return VanillaMTPProposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) raise ValueError(f"unsupported speculative mode: {spec_mode}") diff --git a/lightllm/server/router/model_infer/speculative/proposers/base.py b/lightllm/server/router/model_infer/speculative/proposers/base.py index c4d1ff1d80..6bcc1be004 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/speculative/proposers/base.py @@ -1,72 +1,65 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, List, Optional, Union +from typing import TYPE_CHECKING, Optional, Tuple import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput if TYPE_CHECKING: - from lightllm.server.router.model_infer.speculative.engine import SpecEngine + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend @dataclass class SpecProposal: - """Draft proposal returned by a proposer. - - `token_ids` is the LightLLM service equivalent of DeepSpec's - DraftProposal.verify_input_ids. It contains the target model's freshly - sampled token in column 0, followed by draft candidates: - - token_ids: [verify_batch, draft_step + 1] - - With fixed scheduling, `draft_step == backend.max_draft_step`. Dynamic speculative - scheduling may make it shorter, and the engine pads before scatter. - - `draft_probs` is intentionally narrower than DeepSpec's full - [B, K, vocab] probability tensor. The current dynamic scheduler only needs - the selected-token probability from each draft step: - - draft_probs[i]: [verify_batch] - - `schedule_probs` optionally overrides `draft_probs` for dynamic verify - row selection. It can be a list of per-step vectors or a dense - [verify_batch, draft_step] matrix. DSpark uses confidence-head conditional - acceptance probabilities here; the engine scatters them into the same - per-request buffer and the dynamic selector converts them to prefix - survival probabilities. - - `extra_mem_indexes_cpu` records draft-only KV slots allocated by recurrent - or block proposers. Eagle3 uses these slots for recurrent draft tokens; - DFlash uses them for current-block scratch query/mask K/V: - - extra_mem_indexes_cpu: [slot_count] - + """Candidate tokens and scheduling metadata produced by a proposer. + + `token_ids` has shape `[verify_batch, draft_step + 1]`; column 0 contains + target-model tokens and the remaining columns contain draft candidates. + `schedule_scores`, when present, has shape `[verify_batch, draft_step]`. + Each column contains the proposer-specific score used by dynamic scheduling: + selected-token probability for standard proposers, or confidence-head + probability for DSpark. + `schedule_scores_cpu` is the asynchronous CPU copy consumed by planners + that use proposal scores directly. + `extra_mem_indexes_cpu` tracks temporary KV slots owned by the proposal. """ token_ids: torch.Tensor extra_mem_indexes_cpu: Optional[torch.Tensor] - draft_probs: Optional[List[torch.Tensor]] = None - schedule_probs: Optional[Union[List[torch.Tensor], torch.Tensor]] = None + schedule_scores: Optional[torch.Tensor] = None + schedule_scores_cpu: Optional[torch.Tensor] = None class BaseSpecProposer: """Base class for algorithm-specific draft proposal generation. - A proposer owns the draft-side state transition. The target model gives it + A proposer owns the draft-side state transition. The target model gives it the current target token ids plus captured target hidden features through - SpecEngine.prepare_draft_* methods. The proposer returns candidate ids + the prefill-state and proposal hooks. The proposer returns candidate ids but does not verify acceptance; verification is handled by SpecEngine. """ - def __init__(self, engine: "SpecEngine") -> None: - self.engine = engine - self.backend = engine.backend + def __init__(self, *, backend: "ModeBackend", enable_dynamic_spec: bool) -> None: + self.backend = backend + self.enable_dynamic_spec = bool(enable_dynamic_spec) - @property - def enable_dynamic_spec(self) -> bool: - return self.engine.enable_dynamic_spec + def get_draft_steps(self) -> Tuple[int, ...]: + """Return the draft configurations supported by this proposer.""" + + raise NotImplementedError + + def get_draft_cost_ms( + self, + draft_infer_costs, + req_num: int, + verify_batch_size: int, + draft_step: int, + ) -> float: + """Return the complete draft cost for one ``(N, B, d)`` configuration.""" + + raise NotImplementedError def select_accepted_tail_rows(self, b_req_mtp_start_loc: torch.Tensor, accept_len: torch.Tensor) -> torch.Tensor: return (b_req_mtp_start_loc + accept_len - 1).to(torch.long) @@ -88,15 +81,23 @@ def scatter_selected_step_probs( def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: """Allocate draft-owned temporary KV slots.""" - return self.engine.alloc_extra_mem_indexes(token_count) + token_count = int(token_count) + if token_count == 0: + return torch.empty((0,), dtype=torch.int32, device="cpu") + + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + if g_infer_context.radix_cache is not None: + g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) + return g_infer_context.req_manager.mem_manager.alloc(token_count) - def build_initial_draft_state( + def build_draft_state_from_prefill( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - """Build initial draft KV/state before the first decode verify step. + """Build draft KV/state from target prefill before the first decode verify. Inputs: - `target_model_input`: target prompt ModelInput. Its request order and @@ -113,7 +114,7 @@ def build_initial_draft_state( raise NotImplementedError - def build_initial_draft_state_overlap( + def build_draft_state_from_prefill_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -122,7 +123,7 @@ def build_initial_draft_state_overlap( target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - """Build initial draft state for two overlapped prefill microbatches.""" + """Build draft state from two overlapped target-prefill microbatches.""" raise NotImplementedError diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index fd5128ffc9..a67633d4f2 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -22,21 +22,35 @@ class DFlashProposer(BaseSpecProposer): `SpecProposal.extra_mem_indexes_cpu` """ + def get_draft_steps(self): + return (self.backend.max_draft_step,) + + def get_draft_cost_ms( + self, + draft_infer_costs, + req_num: int, + verify_batch_size: int, + draft_step: int, + ) -> float: + block_size = self.backend.draft_models[0].block_size + # Block drafting first commits all verified rows, then generates one + # complete checkpoint-defined block for each request. + extend_cost_ms = draft_infer_costs.estimate(verify_batch_size) + block_cost_ms = draft_infer_costs.get(req_num * block_size) + return extend_cost_ms + block_cost_ms + @torch.no_grad() - def build_initial_draft_state( + def build_draft_state_from_prefill( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: target_hidden = target_model_output.spec_hidden - assert target_hidden is not None if target_hidden.numel() == 0: return draft_model = self.backend.draft_models[0] - assert target_model_input.input_ids is not None - assert target_model_input.input_ids.shape[0] == target_hidden.shape[0] draft_input = copy.copy(target_model_input) # DFlash consumes target hidden states directly on this prefill path. draft_input.mtp_draft_input_hiddens = target_hidden @@ -52,20 +66,15 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> SpecProposal: - assert 0 <= draft_step <= self.backend.max_draft_step - assert accept_len is not None, "DFlash proposal requires target accept lengths" - num_reqs = int(b_req_mtp_start_loc.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - assert accept_len.shape[0] == num_reqs token_ids = next_token_ids.new_full( (next_token_ids.shape[0], draft_step + 1), fill_value=1, ) token_ids[:, 0] = next_token_ids - assert main_model_output is not None and main_model_output.spec_hidden is not None self.extend_draft_kv_cache( main_model_input=main_model_input, target_hidden=main_model_output.spec_hidden, @@ -95,38 +104,27 @@ def propose_next( flat_token_ids, flat_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) else: flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) - assert flat_token_ids.numel() == num_reqs * block_size block_token_ids = flat_token_ids.reshape(num_reqs, block_size) - # Standard DFlash has one leading bonus row; DeepSpec checkpoints do not. - bonus_rows = block_size - self.backend.max_draft_step - assert bonus_rows in (0, 1), ( - f"DFlash block_size={block_size} must equal mtp_step={self.backend.max_draft_step} " - f"or mtp_step + 1={self.backend.max_draft_step + 1}" - ) - token_ids[selected_rows, 1:] = block_token_ids[:, bonus_rows : bonus_rows + draft_step] + token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] - draft_probs = None + schedule_scores = None if self.enable_dynamic_spec: block_token_probs = flat_token_probs.reshape(num_reqs, block_size) - selected_token_probs = block_token_probs[:, bonus_rows : bonus_rows + draft_step] - draft_probs = [ - self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=selected_token_probs[:, step], - verify_row_count=next_token_ids.shape[0], - ) - for step in range(draft_step) - ] + selected_token_probs = block_token_probs[:, :draft_step] + schedule_scores = self.scatter_selected_step_probs( + selected_rows=selected_rows, + selected_probs=selected_token_probs, + verify_row_count=next_token_ids.shape[0], + ) return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=draft_mem_indexes_cpu, - draft_probs=draft_probs, + schedule_scores=schedule_scores, ) def extend_draft_kv_cache(self, main_model_input: ModelInput, target_hidden: torch.Tensor) -> None: draft_model = self.backend.draft_models[0] batch_size = int(target_hidden.shape[0]) - assert batch_size == main_model_input.b_req_idx.shape[0] draft_kv_input = copy.copy(main_model_input) draft_kv_input.batch_size = batch_size @@ -156,7 +154,6 @@ def extend_draft_kv_cache(self, main_model_input: ModelInput, target_hidden: tor device=target_hidden.device, ) draft_kv_input.b_position_delta = None - draft_kv_input.b_prefill_has_output_cpu = [False for _ in range(batch_size)] draft_kv_input.mtp_draft_input_hiddens = target_hidden draft_model.forward(draft_kv_input) @@ -168,9 +165,7 @@ def build_block_draft_input( num_reqs: int, ): draft_model = self.backend.draft_models[0] - num_reqs = int(num_reqs) block_size = int(draft_model.block_size) - assert selected_rows.shape[0] == num_reqs draft_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * block_size) draft_input_ids = next_token_ids.new_full( @@ -193,7 +188,6 @@ def build_block_draft_input( draft_input.batch_size = draft_input.total_token_num draft_input.max_q_seq_len = 1 draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size - draft_input.max_cache_len = draft_input.max_kv_seq_len draft_input.draft_step = block_size - 1 draft_input.b_req_idx = ( main_model_input.b_req_idx.index_select(0, selected_rows).repeat_interleave(block_size).contiguous() diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/speculative/proposers/dspark.py index a282ee2ca8..fc40151ed4 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dspark.py @@ -3,7 +3,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer @@ -26,26 +26,25 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> SpecProposal: - assert 0 <= draft_step <= self.backend.max_draft_step - assert accept_len is not None, "DSpark proposal requires target accept lengths" - num_reqs = int(b_req_mtp_start_loc.shape[0]) verify_row_count = next_token_ids.shape[0] draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - assert block_size == self.backend.max_draft_step, ( - f"DSpark requires --mtp_step={block_size} for this checkpoint, " f"got {self.backend.max_draft_step}" - ) - assert accept_len.shape[0] == num_reqs - proposal_token_ids = next_token_ids.new_full( (verify_row_count, draft_step + 1), fill_value=1, ) proposal_token_ids[:, 0] = next_token_ids - schedule_probs = [] if self.enable_dynamic_spec else None + schedule_scores = ( + torch.empty( + (verify_row_count, draft_step), + dtype=torch.float32, + device=next_token_ids.device, + ) + if self.enable_dynamic_spec + else None + ) - assert main_model_output is not None and main_model_output.spec_hidden is not None self.extend_draft_kv_cache( main_model_input=main_model_input, target_hidden=main_model_output.spec_hidden, @@ -55,7 +54,7 @@ def propose_next( return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, - schedule_probs=schedule_probs, + schedule_scores=schedule_scores, ) selected_rows = self.select_accepted_tail_rows( @@ -69,20 +68,11 @@ def propose_next( num_reqs=num_reqs, ) draft_model_output = draft_model.forward(draft_input) - assert isinstance(draft_model_output, DSparkModelOutput) - expected_block_rows = num_reqs * block_size - assert draft_model_output.logits.ndim >= 2, "draft logits must have a leading block-row dimension" - assert ( - draft_model_output.logits.shape[0] == expected_block_rows - ), f"draft logits rows must be {expected_block_rows}, got {draft_model_output.logits.shape[0]}" if draft_model_output.draft_token_ids is None: flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) else: flat_token_ids = draft_model_output.draft_token_ids - assert ( - flat_token_ids.numel() == expected_block_rows - ), f"draft token rows must be {expected_block_rows}, got {flat_token_ids.numel()}" block_token_ids = flat_token_ids.reshape(num_reqs, block_size) proposal_token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] @@ -90,23 +80,24 @@ def propose_next( confidence_logits = draft_model_output.confidence_logits if confidence_logits is None: raise RuntimeError("DSpark dynamic verify requires confidence head logits") - assert confidence_logits.ndim == 2, "confidence logits must be [selected_rows, block_size]" - assert ( - confidence_logits.shape[0] == num_reqs - ), f"confidence logits rows must be {num_reqs}, got {confidence_logits.shape[0]}" - assert ( - confidence_logits.shape[1] == block_size - ), f"confidence logits must have {block_size} columns, got {confidence_logits.shape[1]}" # Match the clamp used by the GPU dynamic row selector before it # converts conditional confidence to prefix survival probability. - schedule_probs = self.scatter_selected_step_probs( + schedule_scores = self.scatter_selected_step_probs( selected_rows=selected_rows, selected_probs=confidence_logits[:, :draft_step].sigmoid().clamp(min=0.01, max=0.99), verify_row_count=verify_row_count, ) + schedule_scores_cpu = None + if schedule_scores is not None: + schedule_scores_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="dspark_confidence_probs", + gpu_tensor=schedule_scores, + ) + return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=draft_mem_indexes_cpu, - schedule_probs=schedule_probs, + schedule_scores=schedule_scores, + schedule_scores_cpu=schedule_scores_cpu, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py index 0bb3bfe04b..c150f56eb4 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py @@ -1,7 +1,6 @@ from __future__ import annotations import math -import os import torch @@ -18,26 +17,27 @@ class Eagle3Proposer(RecurrentEagleMTPProposer): from that corrected state. """ - def __init__(self, engine) -> None: - super().__init__(engine) - self._confidence_draft_prune = ( - os.getenv( - # Dynamic Eagle3 already ranks target verify rows by proposal - # confidence. Apply the same budget to deep draft frontiers by - # default; static Eagle3 is explicitly excluded below. - "LIGHTLLM_EAGLE3_CONFIDENCE_DRAFT_PRUNE", - "1", - ).lower() - in {"1", "true", "yes", "on"} - ) - self._draft_prune_safety_factor = max( - 0.0, - float(os.getenv("LIGHTLLM_EAGLE3_DRAFT_PRUNE_SAFETY_FACTOR", "1.10")), - ) - self._draft_prune_min_depth = max( - 2, - int(os.getenv("LIGHTLLM_EAGLE3_DRAFT_PRUNE_MIN_DEPTH", "4")), - ) + _DRAFT_PRUNE_SAFETY_FACTOR = 1.10 + _DRAFT_PRUNE_MIN_DEPTH = 4 + + def get_draft_cost_ms( + self, + draft_infer_costs, + req_num: int, + verify_batch_size: int, + draft_step: int, + ) -> float: + draft_cost_ms = draft_infer_costs.estimate(verify_batch_size) + active_count = req_num + draft_row_budget = max(1, verify_batch_size - req_num) + for step in range(1, draft_step): + active_count = self._get_pruned_active_count( + current_count=active_count, + draft_row_budget=draft_row_budget, + next_depth=step + 1, + ) + draft_cost_ms += draft_infer_costs.get(active_count) + return draft_cost_ms def _get_pruned_active_count( self, @@ -45,14 +45,14 @@ def _get_pruned_active_count( draft_row_budget: int, next_depth: int, ) -> int: - if not self._confidence_draft_prune or next_depth < self._draft_prune_min_depth or current_count <= 1: + if next_depth < self._DRAFT_PRUNE_MIN_DEPTH or current_count <= 1: return current_count # A selected token at depth d consumes all d prefix draft rows from # that request. Therefore at most L/d chains can reach depth d when - # the next target verify has L draft-row slots. Keep a configurable - # safety margin, then retain the highest-survival chain frontiers. - active_count = math.ceil(self._draft_prune_safety_factor * max(1, draft_row_budget) / next_depth) + # the next target verify has L draft-row slots. Keep a small safety + # margin, then retain the highest-survival chain frontiers. + active_count = math.ceil(self._DRAFT_PRUNE_SAFETY_FACTOR * max(1, draft_row_budget) / next_depth) return min(current_count, max(1, active_count)) def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: @@ -67,9 +67,6 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> SpecProposal: - assert 0 <= draft_step <= self.backend.max_draft_step - assert accept_len is not None, "Eagle3 proposal requires target accept lengths" - assert main_model_output is not None and main_model_output.spec_hidden is not None verify_row_count = next_token_ids.shape[0] num_reqs = b_req_mtp_start_loc.shape[0] proposal_token_ids = next_token_ids.new_full( @@ -78,7 +75,15 @@ def propose_next( ) proposal_token_ids[:, 0].copy_(next_token_ids) collect_dynamic_probs = self.enable_dynamic_spec - draft_probs = [] if collect_dynamic_probs else None + schedule_scores = ( + torch.zeros( + (verify_row_count, draft_step), + dtype=torch.float32, + device=next_token_ids.device, + ) + if collect_dynamic_probs + else None + ) target_hidden = main_model_output.spec_hidden @@ -100,7 +105,7 @@ def propose_next( return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, - draft_probs=draft_probs, + schedule_scores=schedule_scores, ) draft_logits = draft_model_output.logits.index_select(0, selected_rows) @@ -108,24 +113,22 @@ def propose_next( draft_next_token_ids, selected_draft_prob = self._gen_argmax_token_ids_and_prob( ModelOutput(logits=draft_logits) ) - draft_prob = self.scatter_selected_step_probs( + schedule_scores[:, 0] = self.scatter_selected_step_probs( selected_rows=selected_rows, selected_probs=selected_draft_prob, verify_row_count=verify_row_count, ) - draft_probs.append(draft_prob) chain_survival = selected_draft_prob.float().clamp(0.01, 0.99) else: draft_next_token_ids = self._gen_argmax_token_ids(ModelOutput(logits=draft_logits)) chain_survival = None - assert draft_model_output.spec_hidden is not None draft_hidden = draft_model_output.spec_hidden.index_select(0, selected_rows) proposal_token_ids[selected_rows, 1] = draft_next_token_ids if draft_step == 1: return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, - draft_probs=draft_probs, + schedule_scores=schedule_scores, ) eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * (draft_step - 1)) @@ -154,7 +157,6 @@ def propose_next( else int(selected_rows.shape[0]) ) if active_count < int(selected_rows.shape[0]): - assert chain_survival is not None keep_rows = torch.topk( chain_survival, k=active_count, @@ -189,22 +191,20 @@ def propose_next( draft_output = draft_model.forward(draft_input) if collect_dynamic_probs: draft_next_token_ids, selected_draft_prob = self._gen_argmax_token_ids_and_prob(draft_output) - draft_prob = self.scatter_selected_step_probs( + schedule_scores[:, step] = self.scatter_selected_step_probs( selected_rows=selected_rows, selected_probs=selected_draft_prob, verify_row_count=verify_row_count, ) - draft_probs.append(draft_prob) chain_survival = chain_survival * selected_draft_prob.float().clamp(0.01, 0.99) else: draft_next_token_ids = self._gen_argmax_token_ids(draft_output) proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids draft_hidden = draft_output.spec_hidden - assert draft_hidden is not None selected_seq_len = selected_seq_len + 1 return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=eagle_mem_indexes_cpu, - draft_probs=draft_probs, + schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py index 4e71fb1679..194b415964 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -11,7 +11,23 @@ class RecurrentEagleMTPProposer(BaseSpecProposer): """Shared draft-state setup for recurrent Eagle MTP proposers.""" - def build_initial_draft_state( + def get_draft_steps(self): + return tuple(range(1, self.backend.max_draft_step + 1)) + + def get_draft_cost_ms( + self, + draft_infer_costs, + req_num: int, + verify_batch_size: int, + draft_step: int, + ) -> float: + # The mandatory extend processes all verified rows and produces the + # first candidate. Later recurrent forwards process one row per request. + extend_cost_ms = draft_infer_costs.estimate(verify_batch_size) + decode_cost_ms = draft_infer_costs.get(req_num) * (draft_step - 1) + return extend_cost_ms + decode_cost_ms + + def build_draft_state_from_prefill( self, target_model_input: ModelInput, target_model_output: ModelOutput, @@ -19,7 +35,6 @@ def build_initial_draft_state( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - assert target_model_output.spec_hidden is not None draft_model_input = prepare_mtp_prefill_inputs( model_input=target_model_input, b_next_token_ids=next_token_ids, @@ -27,7 +42,7 @@ def build_initial_draft_state( ) self.backend.draft_models[0].forward(draft_model_input) - def build_initial_draft_state_overlap( + def build_draft_state_from_prefill_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -38,8 +53,6 @@ def build_initial_draft_state_overlap( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - assert target_model_output0.spec_hidden is not None - assert target_model_output1.spec_hidden is not None draft_model_input0 = prepare_mtp_prefill_inputs( model_input=target_model_input0, b_next_token_ids=next_token_ids0, @@ -70,8 +83,6 @@ def make_verify_extend_input( ) -> ModelInput: new_input = copy.copy(base_input) batch_size = int(input_ids.shape[0]) - assert batch_size == int(draft_hidden.shape[0]) - assert batch_size == int(base_input.b_seq_len.shape[0]) new_input.is_prefill = True new_input.batch_size = batch_size new_input.total_token_num = batch_size @@ -169,7 +180,6 @@ def propose_next_overlap( HOLD rows required to keep the two microbatches shape-compatible). """ - assert 0 <= draft_step <= self.backend.max_draft_step verify_width = self.backend.max_draft_step + 1 inputs = (main_model_input0, main_model_input1) outputs = (main_model_output0, main_model_output1) @@ -184,13 +194,8 @@ def propose_next_overlap( for model_input, model_output, token_ids, real_rows, accept_len in zip( inputs, outputs, next_ids, real_verify_rows, accept_lens ): - assert model_output.spec_hidden is not None - assert model_input.batch_size % verify_width == 0 - assert real_rows % verify_width == 0 - assert token_ids.shape[0] == model_input.batch_size request_capacity = model_input.batch_size // verify_width real_request_num = real_rows // verify_width - assert accept_len.shape[0] == request_capacity starts = torch.arange( 0, model_input.batch_size, @@ -243,7 +248,6 @@ def propose_next_overlap( selected_output = ModelOutput(logits=extend_output.logits.index_select(0, selected)) step_token_ids = self._gen_argmax_token_ids(selected_output) draft_next_token_ids.append(step_token_ids) - assert extend_output.spec_hidden is not None draft_hiddens.append(extend_output.spec_hidden.index_select(0, selected)) selected_seq_lens.append(model_input.b_seq_len.index_select(0, selected) + 1) selected_req_idxs.append(model_input.b_req_idx.index_select(0, selected)) @@ -309,7 +313,6 @@ def propose_next_overlap( step_token_ids = self._gen_argmax_token_ids(step_output) draft_next_token_ids[index] = step_token_ids draft_hiddens[index] = step_output.spec_hidden - assert draft_hiddens[index] is not None selected_seq_lens[index] = selected_seq_lens[index] + 1 real_request_num = real_request_nums[index] if real_request_num > 0: @@ -329,6 +332,108 @@ class EagleMTPProposer(RecurrentEagleMTPProposer): hidden back into the same draft model. """ + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor, + draft_step: int, + ) -> SpecProposal: + """Run the fixed-layout Eagle MTP decode for two DP microbatches. + + DPEP overlap requires every rank to keep the target verify layout + throughout draft decode. Each logical request therefore remains a + contiguous group of ``draft_step + 1`` rows; accepted rows are not + compacted into a one-row recurrent batch on this path. + """ + + verify_width = self.backend.max_draft_step + 1 + model_inputs = (main_model_input0, main_model_input1) + real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) + real_request_nums = tuple(row_count // verify_width for row_count in real_verify_rows) + request_capacities = tuple(model_input.batch_size // verify_width for model_input in model_inputs) + total_real_requests = sum(real_request_nums) + + token_id_steps = [ + torch.cat( + [ + next_token_ids0[:real_verify_rows0], + next_token_ids1[:real_verify_rows1], + ], + dim=0, + ) + ] + if draft_step == 0: + return SpecProposal( + token_ids=token_id_steps[0].unsqueeze(1), + extra_mem_indexes_cpu=None, + ) + + extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(total_real_requests * draft_step) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) + split = real_request_nums[0] * draft_step + microbatch_mem_indexes = ( + extra_mem_indexes[:split], + extra_mem_indexes[split:], + ) + + draft_token_ids = [next_token_ids0, next_token_ids1] + draft_hiddens = [main_model_output0.spec_hidden, main_model_output1.spec_hidden] + draft_model = self.backend.draft_models[0] + hold_mem_index = self.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX + + for step in range(draft_step): + for index, model_input in enumerate(model_inputs): + model_input.input_ids = draft_token_ids[index] + model_input.mtp_draft_input_hiddens = draft_hiddens[index] + + draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) + + for index, (model_input, draft_output) in enumerate(zip(model_inputs, draft_outputs)): + model_input.b_seq_len += 1 + model_input.max_kv_seq_len += 1 + + real_request_num = real_request_nums[index] + mem_start = step * real_request_num + step_mem_indexes = microbatch_mem_indexes[index][mem_start : mem_start + real_request_num] + step_mem_indexes = self._pad_step_mem_indexes( + real_mem_indexes=step_mem_indexes, + request_capacity=request_capacities[index], + hold_mem_index=hold_mem_index, + ) + model_input.mem_indexes = torch.cat( + [ + model_input.mem_indexes.view(-1, verify_width)[:, 1:], + step_mem_indexes.view(-1, 1), + ], + dim=1, + ).view(-1) + + draft_token_ids[index] = self._gen_argmax_token_ids(draft_output) + draft_hiddens[index] = draft_output.spec_hidden + + token_id_steps.append( + torch.cat( + [ + draft_token_ids[0][:real_verify_rows0], + draft_token_ids[1][:real_verify_rows1], + ], + dim=0, + ) + ) + + return SpecProposal( + token_ids=torch.stack(token_id_steps, dim=1), + extra_mem_indexes_cpu=extra_mem_indexes_cpu, + ) + def propose_next( self, main_model_input: ModelInput, @@ -338,9 +443,6 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> SpecProposal: - assert 0 <= draft_step <= self.backend.max_draft_step - assert accept_len is not None - assert main_model_output is not None and main_model_output.spec_hidden is not None verify_row_count = int(next_token_ids.shape[0]) num_reqs = int(b_req_mtp_start_loc.shape[0]) selected_rows = self.select_accepted_tail_rows( @@ -352,7 +454,15 @@ def propose_next( fill_value=1, ) proposal_token_ids[:, 0].copy_(next_token_ids) - draft_probs = [] if self.enable_dynamic_spec else None + schedule_scores = ( + torch.zeros( + (verify_row_count, draft_step), + dtype=torch.float32, + device=next_token_ids.device, + ) + if self.enable_dynamic_spec + else None + ) draft_model = self.backend.draft_models[0] extend_input = self.make_verify_extend_input( @@ -366,31 +476,28 @@ def propose_next( return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, - draft_probs=draft_probs, + schedule_scores=schedule_scores, ) selected_logits = extend_output.logits.index_select(0, selected_rows) selected_output = ModelOutput(logits=selected_logits) if self.enable_dynamic_spec: draft_next_token_ids, selected_prob = self._gen_argmax_token_ids_and_prob(selected_output) - draft_probs.append( - self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=selected_prob, - verify_row_count=verify_row_count, - ) + schedule_scores[:, 0] = self.scatter_selected_step_probs( + selected_rows=selected_rows, + selected_probs=selected_prob, + verify_row_count=verify_row_count, ) else: draft_next_token_ids = self._gen_argmax_token_ids(selected_output) proposal_token_ids[selected_rows, 1] = draft_next_token_ids - assert extend_output.spec_hidden is not None draft_hidden = extend_output.spec_hidden.index_select(0, selected_rows) if draft_step == 1: return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, - draft_probs=draft_probs, + schedule_scores=schedule_scores, ) eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * (draft_step - 1)) @@ -422,22 +529,19 @@ def propose_next( draft_output = draft_model.forward(draft_input) if self.enable_dynamic_spec: draft_next_token_ids, selected_prob = self._gen_argmax_token_ids_and_prob(draft_output) - draft_probs.append( - self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=selected_prob, - verify_row_count=verify_row_count, - ) + schedule_scores[:, step] = self.scatter_selected_step_probs( + selected_rows=selected_rows, + selected_probs=selected_prob, + verify_row_count=verify_row_count, ) else: draft_next_token_ids = self._gen_argmax_token_ids(draft_output) proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids draft_hidden = draft_output.spec_hidden - assert draft_hidden is not None selected_seq_len = selected_seq_len + 1 return SpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=eagle_mem_indexes_cpu, - draft_probs=draft_probs, + schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py index db0fd974be..5ad90cd29c 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py @@ -19,7 +19,20 @@ class VanillaMTPProposer(BaseSpecProposer): draft forward """ - def build_initial_draft_state( + def get_draft_steps(self): + return tuple(range(self.backend.max_draft_step + 1)) + + def get_draft_cost_ms( + self, + draft_infer_costs, + req_num: int, + verify_batch_size: int, + draft_step: int, + ) -> float: + # Each selected MTP module processes the complete target verify batch. + return draft_infer_costs.get(verify_batch_size) * draft_step + + def build_draft_state_from_prefill( self, target_model_input: ModelInput, target_model_output: ModelOutput, @@ -31,7 +44,6 @@ def build_initial_draft_state( source_model_output = target_model_output draft_next_token_ids = next_token_ids for draft_model in self.backend.draft_models: - assert source_model_output.spec_hidden is not None draft_model_input = prepare_mtp_prefill_inputs( model_input=draft_model_input, b_next_token_ids=draft_next_token_ids, @@ -40,7 +52,7 @@ def build_initial_draft_state( source_model_output = draft_model.forward(draft_model_input) draft_next_token_ids = self.backend._gen_argmax_token_ids(source_model_output) - def build_initial_draft_state_overlap( + def build_draft_state_from_prefill_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -59,8 +71,6 @@ def build_initial_draft_state_overlap( draft_next_token_ids1 = next_token_ids1 for draft_model in self.backend.draft_models: - assert source_model_output0.spec_hidden is not None - assert source_model_output1.spec_hidden is not None draft_model_input0 = prepare_mtp_prefill_inputs( model_input=draft_model_input0, b_next_token_ids=draft_next_token_ids0, @@ -87,12 +97,19 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> SpecProposal: - assert 0 <= draft_step <= self.backend.max_draft_step draft_model_input = main_model_input draft_next_token_ids = next_token_ids draft_hidden = main_model_output.spec_hidden if draft_step > 0 else None all_next_token_ids = [next_token_ids] - draft_probs = [] if self.enable_dynamic_spec else None + schedule_scores = ( + torch.empty( + (next_token_ids.shape[0], draft_step), + dtype=torch.float32, + device=next_token_ids.device, + ) + if self.enable_dynamic_spec + else None + ) for step in range(draft_step): draft_model = self.backend.draft_models[step] @@ -100,10 +117,9 @@ def propose_next( draft_model_input.mtp_draft_input_hiddens = draft_hidden draft_model_output = draft_model.forward(draft_model_input) draft_hidden = draft_model_output.spec_hidden - assert draft_hidden is not None if self.enable_dynamic_spec: draft_next_token_ids, draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) - draft_probs.append(draft_prob) + schedule_scores[:, step] = draft_prob else: draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) all_next_token_ids.append(draft_next_token_ids) @@ -111,5 +127,5 @@ def propose_next( return SpecProposal( token_ids=torch.stack(all_next_token_ids, dim=1), extra_mem_indexes_cpu=None, - draft_probs=draft_probs, + schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/speculative/runner.py b/lightllm/server/router/model_infer/speculative/runner.py deleted file mode 100644 index 7ab72b1fc5..0000000000 --- a/lightllm/server/router/model_infer/speculative/runner.py +++ /dev/null @@ -1,207 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass -from typing import TYPE_CHECKING, Callable, List, Optional, Tuple - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.basemodel.triton_kernel.mtp_utils import ( - gen_b_req_mtp_start_loc, - linear_att_spec_state_index_update, -) -from lightllm.server.router.model_infer.infer_batch import g_infer_context -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager - -if TYPE_CHECKING: - from lightllm.server.router.model_infer.speculative.planner import SpecDecodePlan - from lightllm.server.router.model_infer.speculative.engine import SpecEngine - - -@dataclass -class SpecDecodeForwardState: - model_input: ModelInput - original_run_reqs: List - plan: "SpecDecodePlan" - selected_row_mask_cpu: Optional[torch.Tensor] - accepted_index_cpu: torch.Tensor - spec_accept_len_cpu: torch.Tensor - next_token_ids_cpu: torch.Tensor - next_token_logprobs_cpu: torch.Tensor - next_token_ranks_cpu: torch.Tensor - verify_event: torch.cuda.Event - sync_event: torch.cuda.Event - extra_mem_indexes_cpu: Optional[torch.Tensor] - schedule_probs_cpu: Optional[torch.Tensor] - - -@dataclass -class SpecDecodePostState: - next_token_ids: torch.Tensor - next_token_logprobs: torch.Tensor - next_token_ranks: torch.Tensor - spec_accept_len_cpu: torch.Tensor - need_free_mem_indexes: torch.Tensor - - -class SpecDecodeRunner: - def __init__(self, engine: "SpecEngine") -> None: - self.engine = engine - self.backend = engine.backend - - def run_speculative_forward( - self, - model_input: ModelInput, - model_output: ModelOutput, - run_reqs: List, - req_num: int, - plan: "SpecDecodePlan", - selected_row_mask_cpu: Optional[torch.Tensor], - next_token_ids: torch.Tensor, - next_token_logprobs: torch.Tensor, - next_token_ranks: torch.Tensor, - copy_next_token_infos: Callable[ - [torch.Tensor, torch.Tensor, torch.Tensor], - Tuple[torch.Tensor, torch.Tensor, torch.Tensor], - ], - ) -> SpecDecodeForwardState: - engine = self.engine - b_req_mtp_start_loc = gen_b_req_mtp_start_loc(model_input.b_mtp_index, num_reqs=req_num) - accept_len, accepted_index = engine.verify_target_tokens( - new_next_token_ids=next_token_ids, - b_req_idx=model_input.b_req_idx, - b_req_mtp_start_loc=b_req_mtp_start_loc, - ) - if self.backend.is_linear_att_mixed_model: - linear_att_spec_state_index_update( - req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=model_input.b_req_idx, - b_mtp_index=model_input.b_mtp_index, - accepted_index=accepted_index, - verify_width=self.backend.max_draft_step + 1, - ) - accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="accepted_index", - gpu_tensor=accepted_index, - ) - - verify_event = torch.cuda.Event(enable_timing=True) - verify_event.record() - - proposal = engine.propose_next( - main_model_input=model_input, - main_model_output=model_output, - next_token_ids=next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=plan.draft_step, - accept_len=accept_len, - ) - - all_next_token_ids = engine.pad_all_next_token_ids( - token_ids=proposal.token_ids, - draft_step=plan.draft_step, - ) - all_next_token_probs = engine.build_all_next_token_probs( - proposal=proposal, - draft_step=plan.draft_step, - ) - schedule_probs_cpu = ( - g_pin_mem_manager.async_copy_from_gpu_tensor( - key="mtp_schedule_probs", - gpu_tensor=all_next_token_probs, - ) - if all_next_token_probs is not None and engine.needs_schedule_probs_cpu() - else None - ) - - engine.scatter_next_tokens( - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=model_input.b_req_idx, - spec_accept_len=accept_len, - all_next_token_probs=all_next_token_probs, - ) - - next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu = copy_next_token_infos( - next_token_ids, - next_token_logprobs, - next_token_ranks, - ) - - g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( - b_req_idx=model_input.b_req_idx, - next_token_ids=next_token_ids, - mask=accepted_index == 1, - ) - - spec_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="spec_accept_len", - gpu_tensor=accept_len, - ) - - sync_event = torch.cuda.Event() - sync_event.record() - - return SpecDecodeForwardState( - model_input=model_input, - original_run_reqs=run_reqs, - plan=plan, - selected_row_mask_cpu=selected_row_mask_cpu, - accepted_index_cpu=accepted_index_cpu, - spec_accept_len_cpu=spec_accept_len_cpu, - next_token_ids_cpu=next_token_ids_cpu, - next_token_logprobs_cpu=next_token_logprobs_cpu, - next_token_ranks_cpu=next_token_ranks_cpu, - verify_event=verify_event, - sync_event=sync_event, - extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, - schedule_probs_cpu=schedule_probs_cpu, - ) - - def resolve_pre_post_reqs(self, state: SpecDecodeForwardState, decode_reqs: List): - if state.plan.skip_verify_sync: - assert self.engine.enable_dynamic_spec, "skip_verify_sync requires dynamic speculative scheduling" - return decode_reqs, decode_reqs - - state.verify_event.synchronize() - return self.engine.build_decode_req_lists( - original_run_reqs=state.original_run_reqs, - selected_row_mask_cpu=state.selected_row_mask_cpu, - accepted_index_cpu=state.accepted_index_cpu, - ) - - def finish_post(self, state: SpecDecodeForwardState, req_num: int) -> SpecDecodePostState: - state.sync_event.synchronize() - - engine = self.engine - if engine.enable_dynamic_spec: - engine.update_dynamic_accept_stats( - req_num=req_num, - accepted_index_cpu=state.accepted_index_cpu, - spec_accept_len_cpu=state.spec_accept_len_cpu, - dynamic_batch_size=state.plan.dynamic_batch_size, - pre_draft_step=state.plan.pre_draft_step, - ) - if state.schedule_probs_cpu is not None: - engine.update_dynamic_schedule_stats( - req_num=req_num, - schedule_probs_cpu=state.schedule_probs_cpu, - ) - - need_free_mem_indexes = engine.build_decode_free_mem_indexes_cpu( - model_input=state.model_input, - selected_row_mask_cpu=state.selected_row_mask_cpu, - accepted_index_cpu=state.accepted_index_cpu, - ) - if state.extra_mem_indexes_cpu is not None: - need_free_mem_indexes = torch.cat([need_free_mem_indexes, state.extra_mem_indexes_cpu], dim=0) - - select_mask = state.accepted_index_cpu.to(dtype=torch.bool) - return SpecDecodePostState( - next_token_ids=state.next_token_ids_cpu[select_mask], - next_token_logprobs=state.next_token_logprobs_cpu[select_mask], - next_token_ranks=state.next_token_ranks_cpu[select_mask], - spec_accept_len_cpu=state.spec_accept_len_cpu, - need_free_mem_indexes=need_free_mem_indexes, - ) diff --git a/test/test_api/test_gsmk.py b/test/test_api/test_gsmk.py index 2d9ead65b8..ab04d24bed 100644 --- a/test/test_api/test_gsmk.py +++ b/test/test_api/test_gsmk.py @@ -85,34 +85,26 @@ def download_and_cache_file(url: str, filename: Optional[str] = None): def call_generate_lightllm(prompt, temperature, max_tokens, stop=None, url=None): - """Call LightLLM API for text generation.""" + """Call LightLLM through its OpenAI-compatible chat API.""" assert url is not None data = { - "inputs": prompt, - "parameters": { - "temperature": temperature, - "max_new_tokens": max_tokens, - "stop_sequences": stop, - "repetition_penalty": 1.0, - "top_p": 1.0, - "top_k": 1, - }, + "model": "default_model_name", + "messages": [{"role": "user", "content": prompt}], + "temperature": temperature, + "max_tokens": max_tokens, + "stop": stop, + "top_p": 1.0, } - res = requests.post(url, json=data) + res = requests.post(url, json=data, timeout=600) assert res.status_code == 200, f"API request failed with status code {res.status_code}: {res.text}" response_json = res.json() - if "generated_text" not in response_json: - raise ValueError(f"Invalid API response format. Expected 'generated_text' key, got: {response_json.keys()}") - if not isinstance(response_json["generated_text"], list) or len(response_json["generated_text"]) == 0: - raise ValueError( - "Invalid API response format. 'generated_text' should be a non-empty list, " - f"got: {response_json['generated_text']}" - ) - - pred = response_json["generated_text"][0] - return pred + choices = response_json.get("choices") + if not choices: + raise ValueError(f"Invalid chat completion response: {response_json}") + message = choices[0]["message"] + return (message.get("reasoning") or "") + (message.get("content") or "") def get_one_example(lines, i, include_answer): @@ -156,7 +148,9 @@ def parse_args(): parser.add_argument("--port", type=int, default=8000) parser.add_argument("--num-shots", type=int, default=5) parser.add_argument("--num-questions", type=int, default=200) + parser.add_argument("--max-tokens", type=int, default=1024) parser.add_argument("--result-file", type=str, default="result.jsonl") + parser.add_argument("--output-file", type=str, default="tmp_output_lightllm.txt") parser.add_argument("--data-path", type=str, default="test.jsonl") parser.add_argument( "--system-prompt", action="store_true", help="Prepend an 8192-character system prompt to each request" @@ -166,7 +160,7 @@ def parse_args(): def main(args): # LightLLM API URL - url = f"{args.host}:{args.port}/generate" + url = f"{args.host}:{args.port}/v1/chat/completions" # Read data url_data = "https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl" @@ -207,7 +201,7 @@ def get_one_answer(i): answer = call_generate_lightllm( prompt=system_prefix + few_shot_examples + questions[i], temperature=0, - max_tokens=1024, + max_tokens=args.max_tokens, stop=["Question", "Assistant:", "<|separator|>", "Human:", "\n\nQuestion"], url=url, ) @@ -242,7 +236,7 @@ def get_one_answer(i): print(f"Latency: {latency:.3f} s") # Dump results - dump_state_text("tmp_output_lightllm.txt", states) + dump_state_text(args.output_file, states) with open(args.result_file, "a") as fout: value = { diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py index 2fd378b7c5..917fc1cb04 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py @@ -4,61 +4,61 @@ import numpy as np from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import ( - _fwd_kernel_cumprod_probs, + _fwd_kernel_cumprod_scores, sample_dynamic_spec_row_mask, ) -def _reference_cumprod_probs(req_to_next_token_probs, b_req_idx, max_draft_step: int) -> torch.Tensor: - probs = req_to_next_token_probs.clone() +def _reference_cumprod_scores(req_to_next_token_scores, b_req_idx, max_draft_step: int) -> torch.Tensor: + scores = req_to_next_token_scores.clone() req_num = b_req_idx.shape[0] // (max_draft_step + 1) for req_i in range(req_num): req_idx = int(b_req_idx[req_i * (max_draft_step + 1)].item()) - row = probs[req_idx, : max_draft_step + 1].clone() + row = scores[req_idx, : max_draft_step + 1].clone() row[0] = 1.0 row[1:] = torch.clamp(row[1:], min=0.01, max=0.99) - probs[req_idx, : max_draft_step + 1] = torch.cumprod(row, dim=0) - return probs + scores[req_idx, : max_draft_step + 1] = torch.cumprod(row, dim=0) + return scores -def _flat_cumprod_probs( +def _flat_cumprod_scores( b_req_idx: torch.Tensor, - req_to_next_token_probs: torch.Tensor, + req_to_next_token_scores: torch.Tensor, max_draft_step: int, ) -> torch.Tensor: - probs = _reference_cumprod_probs(req_to_next_token_probs, b_req_idx, max_draft_step) + scores = _reference_cumprod_scores(req_to_next_token_scores, b_req_idx, max_draft_step) req_num = b_req_idx.shape[0] // (max_draft_step + 1) all_num = req_num * (max_draft_step + 1) - flat_probs = [] + flat_scores = [] for offset in range(all_num): req_idx = int(b_req_idx[offset].item()) mtp_index = offset % (max_draft_step + 1) - flat_probs.append(probs[req_idx, mtp_index]) - return torch.stack(flat_probs) + flat_scores.append(scores[req_idx, mtp_index]) + return torch.stack(flat_scores) -def _assert_topk_mask(select: torch.Tensor, flat_probs: torch.Tensor, dynamic_batch_size: int) -> None: +def _assert_topk_mask(select: torch.Tensor, flat_scores: torch.Tensor, dynamic_batch_size: int) -> None: k = dynamic_batch_size assert int(select.sum().item()) == k - selected_scores = flat_probs[select.bool()] - unselected_scores = flat_probs[(select == 0).bool()] + selected_scores = flat_scores[select.bool()] + unselected_scores = flat_scores[(select == 0).bool()] if unselected_scores.numel() > 0: assert selected_scores.min() >= unselected_scores.max() - 1e-5 -def _make_batch_probs(req_num: int, max_draft_step: int, rows): +def _make_batch_scores(req_num: int, max_draft_step: int, rows): max_req = req_num - probs = torch.zeros((max_req + 1, 16), dtype=torch.float32, device="cuda") + scores = torch.zeros((max_req + 1, 16), dtype=torch.float32, device="cuda") for req_idx, row in enumerate(rows): - probs[req_idx, : max_draft_step + 1] = torch.tensor(row, dtype=torch.float32, device="cuda") + scores[req_idx, : max_draft_step + 1] = torch.tensor(row, dtype=torch.float32, device="cuda") b_req_idx = torch.arange(req_num, dtype=torch.int32, device="cuda").repeat_interleave(max_draft_step + 1) - return probs, b_req_idx + return scores, b_req_idx @pytest.mark.parametrize("max_draft_step", [1, 3]) -def test_cumprod_probs_kernel(max_draft_step: int): +def test_cumprod_scores_kernel(max_draft_step: int): req_num = 2 - probs, b_req_idx = _make_batch_probs( + scores, b_req_idx = _make_batch_scores( req_num, max_draft_step, rows=[ @@ -66,58 +66,60 @@ def test_cumprod_probs_kernel(max_draft_step: int): [1.0] + [0.2] * max_draft_step, ], ) - probs_clone = probs.clone() - _fwd_kernel_cumprod_probs[(req_num,)]( - req_to_next_token_probs=probs_clone, - req_to_next_token_probs_stride=probs_clone.stride(0), + scores_clone = scores.clone() + _fwd_kernel_cumprod_scores[(req_num,)]( + req_to_next_token_scores=scores_clone, + req_to_next_token_scores_stride=scores_clone.stride(0), b_req_idx=b_req_idx, max_draft_step=max_draft_step, BLOCK_SIZE=triton.next_power_of_2(max_draft_step + 1), num_warps=1, num_stages=1, ) - expected = _reference_cumprod_probs(probs, b_req_idx, max_draft_step) - assert torch.allclose(probs_clone[:, : max_draft_step + 1], expected[:, : max_draft_step + 1], rtol=1e-5, atol=1e-5) + expected = _reference_cumprod_scores(scores, b_req_idx, max_draft_step) + assert torch.allclose( + scores_clone[:, : max_draft_step + 1], expected[:, : max_draft_step + 1], rtol=1e-5, atol=1e-5 + ) -def test_cumprod_probs_clamps_invalid_values(): +def test_cumprod_scores_clamps_invalid_values(): max_draft_step = 2 req_num = 1 - probs, b_req_idx = _make_batch_probs(req_num, max_draft_step, rows=[[1.0, 0.0, 1.5]]) - _fwd_kernel_cumprod_probs[(req_num,)]( - req_to_next_token_probs=probs, - req_to_next_token_probs_stride=probs.stride(0), + scores, b_req_idx = _make_batch_scores(req_num, max_draft_step, rows=[[1.0, 0.0, 1.5]]) + _fwd_kernel_cumprod_scores[(req_num,)]( + req_to_next_token_scores=scores, + req_to_next_token_scores_stride=scores.stride(0), b_req_idx=b_req_idx, max_draft_step=max_draft_step, BLOCK_SIZE=triton.next_power_of_2(max_draft_step + 1), num_warps=1, num_stages=1, ) - row = probs[0, : max_draft_step + 1] + row = scores[0, : max_draft_step + 1] # Index 0 is the always-accepted target sample; clamping applies only to drafts. assert row[0].item() == pytest.approx(1.0) assert row[1].item() == pytest.approx(0.01, rel=1e-4) assert row[2].item() == pytest.approx(0.01 * 0.99, rel=1e-4) -def test_cumprod_probs_clamps_boundary_values(): +def test_cumprod_scores_clamps_boundary_values(): max_draft_step = 3 req_num = 1 - probs, b_req_idx = _make_batch_probs(req_num, max_draft_step, rows=[[1.0, 0.995, 0.005, 0.5]]) - raw_probs = probs.clone() - _fwd_kernel_cumprod_probs[(req_num,)]( - req_to_next_token_probs=probs, - req_to_next_token_probs_stride=probs.stride(0), + scores, b_req_idx = _make_batch_scores(req_num, max_draft_step, rows=[[1.0, 0.995, 0.005, 0.5]]) + raw_scores = scores.clone() + _fwd_kernel_cumprod_scores[(req_num,)]( + req_to_next_token_scores=scores, + req_to_next_token_scores_stride=scores.stride(0), b_req_idx=b_req_idx, max_draft_step=max_draft_step, BLOCK_SIZE=triton.next_power_of_2(max_draft_step + 1), num_warps=1, num_stages=1, ) - expected = _reference_cumprod_probs(raw_probs, b_req_idx, max_draft_step) - row = probs[0, : max_draft_step + 1] + expected = _reference_cumprod_scores(raw_scores, b_req_idx, max_draft_step) + row = scores[0, : max_draft_step + 1] assert torch.allclose(row, expected[0, : max_draft_step + 1], rtol=1e-5, atol=1e-5) - # Draft probabilities 0.995 and 0.005 clamp to 0.99 and 0.01. + # Draft scores 0.995 and 0.005 clamp to 0.99 and 0.01. assert row[0].item() == pytest.approx(1.0) assert row[1].item() == pytest.approx(0.99, rel=1e-4) assert row[2].item() == pytest.approx(0.99 * 0.01, rel=1e-4) @@ -127,7 +129,7 @@ def test_cumprod_probs_clamps_boundary_values(): def test_sample_select_count(): max_draft_step = 3 req_num = 3 - probs, b_req_idx = _make_batch_probs( + scores, b_req_idx = _make_batch_scores( req_num, max_draft_step, rows=[ @@ -141,7 +143,7 @@ def test_sample_select_count(): select = sample_dynamic_spec_row_mask( dynamic_batch_size=dynamic_batch_size, b_req_idx=b_req_idx, - req_to_next_token_probs=probs.clone(), + req_to_next_token_scores=scores.clone(), max_draft_step=max_draft_step, ) assert select.dtype == torch.int32 @@ -152,7 +154,7 @@ def test_sample_select_count(): def test_sample_accepts_numpy_scalar_dynamic_batch_size(): max_draft_step = 3 - probs, b_req_idx = _make_batch_probs( + scores, b_req_idx = _make_batch_scores( 3, max_draft_step, rows=[ @@ -164,7 +166,7 @@ def test_sample_accepts_numpy_scalar_dynamic_batch_size(): select = sample_dynamic_spec_row_mask( dynamic_batch_size=np.int64(8), b_req_idx=b_req_idx, - req_to_next_token_probs=probs, + req_to_next_token_scores=scores, max_draft_step=np.int64(max_draft_step), ) assert int(select.sum().item()) == 8 @@ -172,7 +174,7 @@ def test_sample_accepts_numpy_scalar_dynamic_batch_size(): def test_sample_topk_by_cumprod_score(): max_draft_step = 3 - probs, b_req_idx = _make_batch_probs( + scores, b_req_idx = _make_batch_scores( 3, max_draft_step, rows=[ @@ -181,20 +183,20 @@ def test_sample_topk_by_cumprod_score(): [1.0, 0.99, 0.99, 0.99], ], ) - flat_probs = _flat_cumprod_probs(b_req_idx, probs, max_draft_step) + flat_scores = _flat_cumprod_scores(b_req_idx, scores, max_draft_step) for dynamic_batch_size in [1, 4, 8, 12]: select = sample_dynamic_spec_row_mask( dynamic_batch_size=dynamic_batch_size, b_req_idx=b_req_idx, - req_to_next_token_probs=probs.clone(), + req_to_next_token_scores=scores.clone(), max_draft_step=max_draft_step, ) - _assert_topk_mask(select, flat_probs, dynamic_batch_size) + _assert_topk_mask(select, flat_scores, dynamic_batch_size) def test_sample_picks_highest_cumprod_rows(): max_draft_step = 1 - probs, b_req_idx = _make_batch_probs( + scores, b_req_idx = _make_batch_scores( 2, max_draft_step, rows=[ @@ -202,14 +204,14 @@ def test_sample_picks_highest_cumprod_rows(): [1.0, 0.1], ], ) - flat_probs = _flat_cumprod_probs(b_req_idx, probs, max_draft_step) + flat_scores = _flat_cumprod_scores(b_req_idx, scores, max_draft_step) select = sample_dynamic_spec_row_mask( dynamic_batch_size=2, b_req_idx=b_req_idx, - req_to_next_token_probs=probs.clone(), + req_to_next_token_scores=scores.clone(), max_draft_step=max_draft_step, ) - _assert_topk_mask(select, flat_probs, 2) + _assert_topk_mask(select, flat_scores, 2) # top-2 scores are both 0.99 at mtp_index==0 (req0 and req1 main rows) assert select[0].item() == 1 assert select[2].item() == 1 @@ -217,15 +219,15 @@ def test_sample_picks_highest_cumprod_rows(): def test_sample_single_request(): max_draft_step = 2 - probs, b_req_idx = _make_batch_probs(1, max_draft_step, rows=[[1.0, 0.5, 0.25]]) - flat_probs = _flat_cumprod_probs(b_req_idx, probs, max_draft_step) + scores, b_req_idx = _make_batch_scores(1, max_draft_step, rows=[[1.0, 0.5, 0.25]]) + flat_scores = _flat_cumprod_scores(b_req_idx, scores, max_draft_step) select = sample_dynamic_spec_row_mask( dynamic_batch_size=2, b_req_idx=b_req_idx, - req_to_next_token_probs=probs.clone(), + req_to_next_token_scores=scores.clone(), max_draft_step=max_draft_step, ) - _assert_topk_mask(select, flat_probs, 2) + _assert_topk_mask(select, flat_scores, 2) if __name__ == "__main__": diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py index 036eba833c..c8adbe9825 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -47,7 +47,7 @@ def test_compact_dynamic_spec_model_input(monkeypatch): multimodal_params=[{"row": i, "images": [], "audios": []} for i in range(12)], mtp_draft_input_hiddens=(torch.arange(12 * 5, dtype=torch.float32, device="cuda").reshape(12, 5) + 0.5), ) - req_to_next_token_probs = torch.tensor( + req_to_next_token_scores = torch.tensor( [ [1.0, 0.95, 0.90, 0.10, 0.0, 0.0], [1.0, 0.20, 0.80, 0.80, 0.0, 0.0], @@ -61,7 +61,7 @@ def test_compact_dynamic_spec_model_input(monkeypatch): model_input=model_input, req_num=3, dynamic_batch_size=8, - req_to_next_token_probs=req_to_next_token_probs, + req_to_next_token_scores=req_to_next_token_scores, ) torch.cuda.synchronize() @@ -150,12 +150,24 @@ def test_mtp_verify_scatter_and_start_locations(): dtype=torch.int64, device="cuda", ) + req_to_next_token_scores = torch.full((3, 5), -1.0, dtype=torch.float32, device="cuda") + schedule_scores = torch.tensor( + [[0.9, 0.8], [0.8, 0.7], [0.7, 0.6], [0.6, 0.5], [0.5, 0.4]], + dtype=torch.float32, + device="cuda", + ) spec_accept_len, accepted_index = mtp_utils.mtp_verify( req_to_next_token_ids, b_req_mtp_start_loc, new_next_token_ids, b_req_idx ) mtp_utils.mtp_scatter_next_token_ids( - req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, spec_accept_len + req_to_next_token_ids=req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=b_req_idx, + spec_accept_len=spec_accept_len, + req_to_next_token_scores=req_to_next_token_scores, + schedule_scores=schedule_scores, ) torch.cuda.synchronize() @@ -165,7 +177,33 @@ def test_mtp_verify_scatter_and_start_locations(): assert torch.equal( req_to_next_token_ids.cpu(), torch.tensor( - [[1, 2, 3, -1, -1], [1, 2, 0, -1, -1], [7, 8, 9, 4, 5]], + [[1, 2, 3, 1, 1], [1, 2, 0, -1, -1], [7, 8, 9, 1, 1]], dtype=torch.int64, ), ) + torch.testing.assert_close( + req_to_next_token_scores.cpu(), + torch.tensor( + [[1.0, 0.9, 0.8, 0.0, 0.0], [-1.0, -1.0, -1.0, -1.0, -1.0], [1.0, 0.7, 0.6, 0.0, 0.0]], + dtype=torch.float32, + ), + ) + + +def test_mtp_scatter_handles_zero_draft_step(): + req_to_next_token_ids = torch.full((1, 4), -1, dtype=torch.int64, device="cuda") + req_to_next_token_scores = torch.full((1, 4), -1.0, dtype=torch.float32, device="cuda") + + mtp_utils.mtp_scatter_next_token_ids( + req_to_next_token_ids=req_to_next_token_ids, + b_req_mtp_start_loc=torch.tensor([0], dtype=torch.int32, device="cuda"), + all_next_token_ids=torch.tensor([[42]], dtype=torch.int64, device="cuda"), + b_req_idx=torch.tensor([0], dtype=torch.int32, device="cuda"), + spec_accept_len=torch.tensor([1], dtype=torch.int32, device="cuda"), + req_to_next_token_scores=req_to_next_token_scores, + schedule_scores=torch.empty((1, 0), dtype=torch.float32, device="cuda"), + ) + torch.cuda.synchronize() + + assert req_to_next_token_ids.cpu().tolist() == [[42, 1, 1, 1]] + assert req_to_next_token_scores.cpu().tolist() == [[1.0, 0.0, 0.0, 0.0]] diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index b01108c95c..2b67c64d05 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -5,6 +5,7 @@ from lightllm.common.basemodel import batch_objs from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel +from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel from lightllm.models.qwen3_dspark import model_output as dspark_model_output from lightllm.models.qwen3_dspark.layer_weights import pre_and_post_layer_weight as dspark_pre_post_weight from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer @@ -61,8 +62,16 @@ def test_dspark_no_ref_conversion_dispatches_to_dspark_fields(monkeypatch): assert all(converted != original for converted, original in zip(converted_ptrs, original_ptrs)) -def test_confidence_head_does_not_inherit_model_quantization(monkeypatch): - captured_kwargs = {} +@pytest.mark.parametrize( + "head_config", + [ + {"enable_confidence_head": True}, + {"markov_rank": 4, "markov_head_type": "gated"}, + {"markov_rank": 4, "markov_head_type": "rnn"}, + ], +) +def test_biased_dspark_heads_do_not_inherit_model_quantization(monkeypatch, head_config): + captured_kwargs = [] def init_base_weight(self, data_type, network_config, quant_cfg): self.data_type_ = data_type @@ -70,28 +79,44 @@ def init_base_weight(self, data_type, network_config, quant_cfg): class RecordingROWMMWeight: def __init__(self, **kwargs): - captured_kwargs.update(kwargs) + captured_kwargs.append(kwargs) + + class StubWeight: + def __init__(self, **kwargs): + pass monkeypatch.setattr( dspark_pre_post_weight.Qwen3DFlashPreAndPostLayerWeight, "__init__", init_base_weight, ) + monkeypatch.setattr(dspark_pre_post_weight, "EmbeddingWeight", StubWeight) + monkeypatch.setattr(dspark_pre_post_weight, "LMHeadWeight", StubWeight) monkeypatch.setattr(dspark_pre_post_weight, "ROWMMWeight", RecordingROWMMWeight) quant_method = object() quant_cfg = SimpleNamespace(get_quant_method=lambda *_: quant_method) + network_config = { + "hidden_size": 16, + "vocab_size": 32, + **head_config, + } dspark_pre_post_weight.Qwen3DSparkPreAndPostLayerWeight( data_type=torch.bfloat16, - network_config={ - "hidden_size": 16, - "vocab_size": 32, - "enable_confidence_head": True, - }, + network_config=network_config, quant_cfg=quant_cfg, ) - assert captured_kwargs["quant_method"] is None + assert captured_kwargs + assert all(kwargs["quant_method"] is None for kwargs in captured_kwargs) + + +def test_fixed_dspark_does_not_require_confidence_head(monkeypatch): + monkeypatch.setattr(Qwen3DFlashModel, "_verify_params", lambda self: None) + model = Qwen3DSparkModel.__new__(Qwen3DSparkModel) + model.config = {"enable_confidence_head": False} + + model._verify_params() def test_qwen35_dspark_uses_training_rope_layout(monkeypatch): @@ -172,3 +197,77 @@ def __call__(self, input, alloc_func): expected_tokens = torch.stack(expected_tokens, dim=1) torch.testing.assert_close(sampled_tokens, expected_tokens) + + +def test_vanilla_markov_tp4_sampling_matches_full_logits(monkeypatch): + torch.manual_seed(0) + post_infer = Qwen3DSparkPostLayerInfer.__new__(Qwen3DSparkPostLayerInfer) + post_infer.block_size_ = 3 + post_infer.markov_rank_ = 4 + post_infer.markov_head_type_ = "vanilla" + post_infer.tp_world_size_ = 4 + post_infer.alloc_tensor = torch.empty + + vocab_size = 11 + request_count = 2 + split_indexes = torch.linspace(0, vocab_size, post_infer.tp_world_size_ + 1, dtype=torch.int64) + tp_rank = 2 + local_start = int(split_indexes[tp_rank]) + local_end = int(split_indexes[tp_rank + 1]) + + base_logits = torch.randn(request_count, post_infer.block_size_, vocab_size) + markov_w1 = torch.randn(vocab_size, post_infer.markov_rank_) + markov_w2 = torch.randn(vocab_size, post_infer.markov_rank_) + anchor_token_ids = torch.tensor([1, 7]) + + class MarkovEmbedding: + weight = markov_w1 + + class MarkovLMHead: + weight = markov_w2[local_start:local_end] + tp_vocab_start_id = local_start + + layer_weight = SimpleNamespace( + markov_w1_weight_=MarkovEmbedding(), + markov_w2_weight_=MarkovLMHead(), + ) + local_logits = base_logits[:, :, local_start:local_end].reshape(-1, local_end - local_start).T.contiguous() + + expected_tokens = [] + prev_token_ids = anchor_token_ids + for step_idx in range(post_infer.block_size_): + scores = base_logits[:, step_idx] + torch.nn.functional.linear(markov_w1[prev_token_ids], markov_w2) + prev_token_ids = torch.argmax(scores, dim=-1) + expected_tokens.append(prev_token_ids) + expected_tokens = torch.stack(expected_tokens, dim=1) + + step = 0 + + def gather_tp_winners(output, local_winners, group, async_op): + nonlocal step + prev_tokens = anchor_token_ids if step == 0 else expected_tokens[:, step - 1] + scores = base_logits[:, step] + torch.nn.functional.linear(markov_w1[prev_tokens], markov_w2) + winners = [] + for rank in range(post_infer.tp_world_size_): + start = int(split_indexes[rank]) + end = int(split_indexes[rank + 1]) + values, indexes = scores[:, start:end].max(dim=-1) + winners.append(torch.stack((values, (indexes + start).float()), dim=-1)) + torch.testing.assert_close(local_winners, winners[tp_rank]) + output.copy_(torch.stack(winners).reshape_as(output)) + step += 1 + + monkeypatch.setattr( + "lightllm.models.qwen3_dspark.layer_infer.post_layer_infer.all_gather_into_tensor", + gather_tp_winners, + ) + + sampled_tokens = post_infer._sample_markov( + local_logits=local_logits, + block_hidden=torch.empty(request_count, post_infer.block_size_, 0), + infer_state=SimpleNamespace(dist_group=None), + anchor_token_ids=anchor_token_ids, + layer_weight=layer_weight, + ) + + torch.testing.assert_close(sampled_tokens, expected_tokens) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py index c9dd507089..e9118b71bd 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py @@ -33,6 +33,17 @@ def scatter_next_tokens(self, **kwargs): self.scatter_args = kwargs +def test_padded_token_ids_support_empty_dp_rank(): + padded_token_ids = DPChunkedPrefillBackend._build_padded_next_token_ids( + token_ids=None, + batch_size=4, + copy_len=0, + device=torch.device("cpu"), + ) + + assert torch.equal(padded_token_ids, torch.zeros(4, dtype=torch.int64)) + + def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(): backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) backend.max_draft_step = 7 @@ -64,7 +75,7 @@ def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(): assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) -def test_dp_overlap_eagle_keeps_target_padding_out_of_recurrent_decode(): +def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(): backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) backend.max_draft_step = 7 backend.spec_engine = _RecordingSpecEngine() diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py deleted file mode 100644 index 149a3cccca..0000000000 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_post_process.py +++ /dev/null @@ -1,50 +0,0 @@ -import pytest -import torch - -if not torch.cuda.is_available(): - pytest.skip("requires CUDA", allow_module_level=True) - -from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import trim_post_sample_tensors - - -def test_trim_post_sample_tensors(): - selected = torch.tensor([1, 1, 1, 0, 1, 0, 0, 0, 1, 1, 1, 1], dtype=torch.int32, device="cuda") - selected_rows = torch.where(selected.cpu() == 1)[0] - dynamic_batch_size = int(selected.sum().item()) - - b_req_idx = torch.arange(12, dtype=torch.int32, device="cuda") + 10 - b_temperatures = torch.arange(12, dtype=torch.float32, device="cuda") + 0.5 - b_top_ps = torch.arange(12, dtype=torch.float32, device="cuda") / 100 + 0.8 - b_top_ks = torch.arange(12, dtype=torch.int32, device="cuda") + 100 - b_length_penalty_param = torch.arange(12, dtype=torch.int32, device="cuda") + 200 - b_mask_eos_reqs = torch.tensor( - [True, False, True, False, True, False, True, False, True, False, True, False], - dtype=torch.bool, - device="cuda", - ) - - ( - out_b_req_idx, - out_b_temperatures, - out_b_top_ps, - out_b_top_ks, - out_b_length_penalty_param, - out_b_mask_eos_reqs, - ) = trim_post_sample_tensors( - dynamic_batch_size=dynamic_batch_size, - selected_row_mask=selected, - b_req_idx=b_req_idx, - b_temperatures=b_temperatures, - b_top_ps=b_top_ps, - b_top_ks=b_top_ks, - b_length_penalty_param=b_length_penalty_param, - b_mask_eos_reqs=b_mask_eos_reqs, - ) - torch.cuda.synchronize() - - assert torch.equal(out_b_req_idx.cpu(), b_req_idx.cpu()[selected_rows]) - assert torch.equal(out_b_temperatures.cpu(), b_temperatures.cpu()[selected_rows]) - assert torch.equal(out_b_top_ps.cpu(), b_top_ps.cpu()[selected_rows]) - assert torch.equal(out_b_top_ks.cpu(), b_top_ks.cpu()[selected_rows]) - assert torch.equal(out_b_length_penalty_param.cpu(), b_length_penalty_param.cpu()[selected_rows]) - assert torch.equal(out_b_mask_eos_reqs.cpu(), b_mask_eos_reqs.cpu()[selected_rows]) diff --git a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py index da082ad018..b30c01dfb3 100644 --- a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py @@ -39,15 +39,13 @@ def _target_input(batch_size): b_seq_len=torch.arange(batch_size, dtype=torch.int32) + 4, b_req_idx=torch.arange(batch_size, dtype=torch.int32), b_position_delta=torch.zeros(batch_size, dtype=torch.int32), + mem_indexes=torch.arange(batch_size, dtype=torch.int32), max_kv_seq_len=16, max_cache_len=16, ) -def test_overlap_eagle_extends_verify_rows_then_decodes_logical_batch(monkeypatch): - # Keep this topology test CPU-only; production allocates the same tensor - # through the shared CUDA KV manager. - monkeypatch.setattr(torch.Tensor, "cuda", lambda self, **_: self) +def test_overlap_eagle_keeps_fixed_verify_layout(): draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -59,12 +57,8 @@ def test_overlap_eagle_extends_verify_rows_then_decodes_logical_batch(monkeypatc ), _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) - engine = SimpleNamespace( - backend=backend, - enable_dynamic_spec=False, - alloc_extra_mem_indexes=lambda token_count: torch.arange(token_count, dtype=torch.int32), - ) - proposer = EagleMTPProposer(engine) + proposer = EagleMTPProposer(backend=backend, enable_dynamic_spec=False) + proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) model_input0 = _target_input(batch_size=6) model_input1 = _target_input(batch_size=6) @@ -82,11 +76,16 @@ def test_overlap_eagle_extends_verify_rows_then_decodes_logical_batch(monkeypatc draft_step=2, ) - assert draft_model.extend_batch_sizes == (6, 6) - assert draft_model.decode_batch_sizes == [(2, 2)] + assert draft_model.extend_batch_sizes is None + assert draft_model.decode_batch_sizes == [(6, 6), (6, 6)] assert proposal.token_ids.shape == (9, 3) assert torch.equal(proposal.token_ids[:, 0], torch.tensor([0, 1, 2, 10, 11, 12, 13, 14, 15])) - assert proposal.extra_mem_indexes_cpu.shape == (3,) + expected_draft_tokens = torch.tensor([0, 1, 2, 0, 1, 2, 3, 4, 5]) + assert torch.equal(proposal.token_ids[:, 1], expected_draft_tokens) + assert torch.equal(proposal.token_ids[:, 2], expected_draft_tokens) + assert torch.equal(proposal.extra_mem_indexes_cpu, torch.arange(6, dtype=torch.int32)) + assert torch.equal(model_input0.mem_indexes, torch.tensor([2, 0, 1, 5, 99, 99], dtype=torch.int32)) + assert torch.equal(model_input1.mem_indexes, torch.tensor([2, 2, 4, 5, 3, 5], dtype=torch.int32)) def test_eagle3_maps_draft_token_ids_in_proposer(): diff --git a/unit_tests/server/router/model_infer/speculative/test_planner.py b/unit_tests/server/router/model_infer/speculative/test_planner.py index 58f0bb9ab2..70d110ee77 100644 --- a/unit_tests/server/router/model_infer/speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/speculative/test_planner.py @@ -1,12 +1,72 @@ -import random +from types import SimpleNamespace import numpy as np +import torch +from lightllm.server.router.model_infer.speculative.engine import SpecEngine from lightllm.server.router.model_infer.speculative.planner import ( - DSparkDynamicSpecPlanner, - DynamicSpecPlanner, + DSparkPlanner, FixedSpecPlanner, + LightSpecPlanner, + SpecDecodePlan, ) +from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer +from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer +from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer +from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer +from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer + + +def build_draft_cost_provider(proposer_class, max_draft_step: int = 3, block_size: int = 3): + backend = SimpleNamespace( + max_draft_step=max_draft_step, + draft_models=[SimpleNamespace(block_size=block_size)], + ) + return proposer_class(backend=backend, enable_dynamic_spec=True) + + +def build_lightspec_planner( + max_draft_step: int = 3, + proposer_class=VanillaMTPProposer, + block_size: int = 3, +): + return LightSpecPlanner( + max_draft_step=max_draft_step, + draft_cost_provider=build_draft_cost_provider( + proposer_class=proposer_class, + max_draft_step=max_draft_step, + block_size=block_size, + ), + ) + + +def build_dspark_planner(max_draft_step: int = 3, block_size: int = 3): + return DSparkPlanner( + max_draft_step=max_draft_step, + draft_cost_provider=build_draft_cost_provider( + proposer_class=DSparkProposer, + max_draft_step=max_draft_step, + block_size=block_size, + ), + ) + + +def build_planner(spec_mode: str, enable_dynamic_spec: bool = True): + engine = SpecEngine.__new__(SpecEngine) + engine.spec_mode = spec_mode + engine.enable_dynamic_spec = enable_dynamic_spec + engine.backend = SimpleNamespace( + max_draft_step=3, + draft_models=[SimpleNamespace(block_size=3)], + ) + proposer_class = { + "dspark": DSparkProposer, + "dflash": DFlashProposer, + "eagle3": Eagle3Proposer, + }.get(spec_mode, VanillaMTPProposer) + engine.proposer = build_draft_cost_provider(proposer_class) + return engine._build_decode_planner() def test_fixed_planner_returns_static_plan(): @@ -18,42 +78,400 @@ def test_fixed_planner_returns_static_plan(): assert not plan.skip_verify_sync -def test_dynamic_planner_stays_full_width_until_costs_are_profiled(): - plan = DynamicSpecPlanner(max_draft_step=3, use_random_mode=False).plan(req_num=2, original_batch_size=8) +def test_engine_routes_only_dspark_to_the_confidence_planner(): + assert isinstance(build_planner("eagle3", enable_dynamic_spec=False), FixedSpecPlanner) + assert isinstance(build_planner("dspark"), DSparkPlanner) + + dflash_planner = build_planner("dflash") + assert isinstance(dflash_planner, LightSpecPlanner) + assert dflash_planner.draft_steps == (3,) + + eagle_planner = build_planner("eagle3") + assert isinstance(eagle_planner, LightSpecPlanner) + assert eagle_planner.draft_steps == (1, 2, 3) + + +def test_dynamic_plan_filters_selected_rows(): + plan = SpecDecodePlan(dynamic_batch_size=2, draft_step=3, pre_draft_step=3) + reqs = ["req0", "req0", "req1", "req1"] + + selected_reqs = plan.filter_reqs( + reqs=reqs, + selected_row_mask_cpu=torch.tensor([1, 0, 1, 0], dtype=torch.int32), + ) + + assert selected_reqs == ["req0", "req1"] + + +def test_dynamic_decode_frees_unselected_and_rejected_rows(): + freed = [] + engine = SpecEngine.__new__(SpecEngine) + engine.backend = SimpleNamespace( + model=SimpleNamespace( + req_manager=SimpleNamespace(mem_manager=SimpleNamespace(free=lambda indexes: freed.append(indexes.clone()))) + ) + ) + model_input = SimpleNamespace(mem_indexes_cpu=torch.tensor([10, 11, 12, 13])) + + engine.free_unused_decode_mem( + model_input=model_input, + selected_row_mask_cpu=torch.tensor([1, 0, 1, 0], dtype=torch.bool), + accepted_index_cpu=torch.tensor([1, 0], dtype=torch.int32), + extra_mem_indexes_cpu=torch.tensor([20]), + ) + + assert len(freed) == 1 + assert freed[0].tolist() == [11, 12, 13, 20] + + +def test_engine_records_request_spec_metrics_in_one_pass(): + class MetricReq: + def __init__(self, req_idx: int, mtp_step: int): + self.req_idx = req_idx + self.mtp_step = mtp_step + self.accepted = 0 + self.verified = 0 + self.verify_steps = 0 + + def update_spec_accepted_token_num(self, accept_token_num: int): + self.accepted += accept_token_num + + def update_spec_verify_token_num(self, verify_token_num: int): + self.verified += verify_token_num + + def update_spec_verify_step_num(self, verify_step_num: int): + self.verify_steps += verify_step_num + + engine = SpecEngine.__new__(SpecEngine) + engine.backend = SimpleNamespace(is_master_in_dp=True) + req0 = MetricReq(req_idx=10, mtp_step=3) + req1 = MetricReq(req_idx=11, mtp_step=3) + + engine.record_request_spec_metrics( + decode_reqs=[req0, req1], + accept_lengths_cpu=torch.tensor([2, 1], dtype=torch.int32), + verified_row_reqs=[req0, req0, req1], + ) + + assert (req0.accepted, req0.verified, req0.verify_steps) == (1, 2, 1) + assert (req1.accepted, req1.verified, req1.verify_steps) == (0, 1, 1) + + fixed_req = MetricReq(req_idx=12, mtp_step=3) + engine.record_request_spec_metrics( + decode_reqs=[fixed_req], + accept_lengths_cpu=torch.tensor([3], dtype=torch.int32), + ) + assert (fixed_req.accepted, fixed_req.verified, fixed_req.verify_steps) == (2, 4, 1) + + +def test_lightspec_stays_full_width_until_costs_are_profiled(): + plan = build_lightspec_planner().plan( + req_num=2, + original_batch_size=8, + proposal_req_num=2, + ) assert plan.dynamic_batch_size == 8 assert plan.draft_step == plan.pre_draft_step == 3 -def test_planner_does_not_reset_process_global_random_state(): - random.seed(2027) - expected = random.random() - random.seed(2027) +def test_lightspec_collects_full_width_progress_before_adapting(): + planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + for batch_size in (2, 4, 8): + planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=False) + planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) + + plan = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + + assert plan.dynamic_batch_size == 8 + assert plan.draft_step == plan.pre_draft_step == 3 + - DynamicSpecPlanner(max_draft_step=3) +def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): + planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 10.0)): + planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) + for batch_size, draft_cost in ((2, 0.1), (4, 0.1), (8, 0.2)): + planner.update_infer_cost(batch_size, draft_cost, is_draft_model=True) + for _ in range(8): + planner.update_verified_batch( + accept_lengths=[4, 4], + req_num=2, + dynamic_batch_size=8, + verified_draft_step=3, + ) - assert random.random() == expected + plan = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + + assert plan.dynamic_batch_size == 4 + assert plan.pre_draft_step == 3 + assert plan.draft_step == 1 + + +def test_lightspec_compacts_block_verify_without_changing_draft_shape(): + planner = build_lightspec_planner( + max_draft_step=7, + proposer_class=DFlashProposer, + block_size=7, + ) + for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 3.0), (16, 8.0)): + planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) + planner.update_infer_cost(batch_size=14, infer_cost_ms=0.5, is_draft_model=True) + for _ in range(8): + planner.update_verified_batch( + accept_lengths=[3, 3], + req_num=2, + dynamic_batch_size=16, + verified_draft_step=7, + ) + + plan = planner.plan(req_num=2, original_batch_size=16, proposal_req_num=2) + + assert plan.dynamic_batch_size < 16 + assert plan.draft_step == plan.pre_draft_step == 7 + + +def test_lightspec_bounds_verify_to_existing_proposals(): + planner = build_lightspec_planner() + + cold_start = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=0) + mixed_batch = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=1) + ready_batch = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + + assert cold_start.dynamic_batch_size == 2 + assert mixed_batch.dynamic_batch_size == 5 + assert ready_batch.dynamic_batch_size == 8 + assert not cold_start.record_progress + assert not mixed_batch.record_progress + assert ready_batch.record_progress + + +def test_engine_counts_requests_with_a_previous_proposal(): + engine = SpecEngine.__new__(SpecEngine) + engine.planner = build_lightspec_planner() + model_input = SimpleNamespace(batch_size=8) + + mixed_plan = engine.plan_decode( + model_input=model_input, + decode_reqs=[SimpleNamespace(cur_output_len=1), SimpleNamespace(cur_output_len=2)], + ) + ready_plan = engine.plan_decode( + model_input=model_input, + decode_reqs=[SimpleNamespace(cur_output_len=2), SimpleNamespace(cur_output_len=2)], + ) + + assert mixed_plan.dynamic_batch_size == 5 + assert not mixed_plan.record_progress + assert ready_plan.dynamic_batch_size == 8 + assert ready_plan.record_progress + + +def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): + planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + for batch_size in (2, 4, 8): + planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=False) + planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) + planner.update_verified_batch( + accept_lengths=[2, 2], + req_num=2, + dynamic_batch_size=8, + verified_draft_step=3, + ) + + plan = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + + assert planner.draft_steps == (1, 2, 3) + assert plan.draft_step >= 1 + + +def test_vanilla_cost_provider_prices_each_selected_mtp_module(): + planner = build_lightspec_planner() + planner.update_infer_cost(batch_size=8, infer_cost_ms=0.25, is_draft_model=True) + + draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=planner.draft_infer_costs, + req_num=2, + verify_batch_size=8, + draft_step=3, + ) + + assert draft_cost_ms == 0.75 + + +def test_eagle_cost_provider_prices_extend_and_decode_separately(): + planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + planner.update_infer_cost(batch_size=2, infer_cost_ms=0.25, is_draft_model=True) + + draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=planner.draft_infer_costs, + req_num=2, + verify_batch_size=8, + draft_step=3, + ) + + assert draft_cost_ms == 1.5 + + +def test_eagle3_cost_provider_accounts_for_pruned_recurrent_rows(): + planner = build_lightspec_planner( + max_draft_step=7, + proposer_class=Eagle3Proposer, + ) + for batch_size in (2, 4, 8, 16): + planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) + + draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=planner.draft_infer_costs, + req_num=8, + verify_batch_size=16, + draft_step=7, + ) + + assert draft_cost_ms == 42.0 + + +def test_block_cost_provider_prices_commit_and_complete_block(): + planner = build_lightspec_planner( + max_draft_step=7, + proposer_class=DFlashProposer, + block_size=7, + ) + planner.update_infer_cost(batch_size=14, infer_cost_ms=0.7, is_draft_model=True) + + draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=planner.draft_infer_costs, + req_num=2, + verify_batch_size=16, + draft_step=7, + ) + + assert np.isclose(draft_cost_ms, 1.5) + + +def test_lightspec_records_one_batch_observation_per_configuration(): + planner = build_lightspec_planner() + + planner.update_verified_batch( + accept_lengths=[2, 2], + req_num=2, + dynamic_batch_size=4, + verified_draft_step=1, + ) + planner.update_verified_batch( + accept_lengths=[1, 1], + req_num=2, + dynamic_batch_size=4, + verified_draft_step=3, + ) + + assert planner.progress_ema_by_config[(2, 4, 1)].get() == 1.0 + assert np.isclose(planner.progress_ema_by_config[(2, 4, 3)].get(), 0.95) + assert planner.progress_ema_by_config[(2, 4, 1)].get_count() == 1 + assert planner.progress_ema_by_config[(2, 4, 3)].get_count() == 1 + + +def test_lightspec_high_concurrency_does_not_multiply_ema_updates(): + planner = build_lightspec_planner() + + planner.update_verified_batch( + accept_lengths=[1] * 128, + req_num=128, + dynamic_batch_size=128, + verified_draft_step=1, + ) + + assert planner.progress_ema_by_config[(128, 128, 1)].get_count() == 1 + + +def test_lightspec_transfers_committed_progress_to_a_nearby_verify_shape(): + planner = build_lightspec_planner() + planner.update_verified_batch( + accept_lengths=[2, 1], + req_num=2, + dynamic_batch_size=8, + verified_draft_step=3, + ) + + assert planner._estimate_progress(req_num=2, dynamic_batch_size=4, draft_step=3) == 0.75 + assert planner._estimate_progress(req_num=2, dynamic_batch_size=2, draft_step=3) == 1.0 + assert planner._estimate_progress(req_num=2, dynamic_batch_size=4, draft_step=2) == 0.75 + + +def test_engine_skips_feedback_for_a_mixed_proposal_batch(): + engine = SpecEngine.__new__(SpecEngine) + engine.enable_dynamic_spec = True + engine.planner = build_lightspec_planner() + plan = SpecDecodePlan( + dynamic_batch_size=5, + draft_step=3, + pre_draft_step=3, + record_progress=False, + ) + + engine.update_planner_feedback( + plan=plan, + proposal=SpecProposal( + token_ids=torch.empty((0,), dtype=torch.int64), + extra_mem_indexes_cpu=None, + ), + req_num=2, + accept_lengths_cpu=torch.tensor([1, 2], dtype=torch.int32), + ) + + assert not engine.planner.progress_ema_by_config def test_dspark_applies_confidence_capacity_after_two_step_delay(): - planner = DSparkDynamicSpecPlanner(max_draft_step=3) - planner.update_infer_cost(batch_size=2, infer_cost_ms=1.0, is_draft_model=False) - planner.update_infer_cost(batch_size=4, infer_cost_ms=1.1, is_draft_model=False) - planner.update_infer_cost(batch_size=8, infer_cost_ms=10.0, is_draft_model=False) - schedule_probs = np.asarray([[1.0, 0.9, 0.9, 0.9]] * 2, dtype=np.float64) + planner = build_dspark_planner() + for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 10.0)): + planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) + planner.update_infer_cost(batch_size=6, infer_cost_ms=0.5, is_draft_model=True) + confidence_probs = np.asarray([[0.9, 0.9, 0.9]] * 2, dtype=np.float64) + plan = SpecDecodePlan(dynamic_batch_size=8, draft_step=3, pre_draft_step=3) + proposal = SpecProposal( + token_ids=torch.empty((0,), dtype=torch.int64), + extra_mem_indexes_cpu=None, + schedule_scores_cpu=torch.from_numpy(confidence_probs), + ) + engine = SpecEngine.__new__(SpecEngine) + engine.enable_dynamic_spec = True + engine.planner = planner - planner.update_predicted_schedule_probs(schedule_probs=schedule_probs, req_num=2) + engine.update_planner_feedback( + plan=plan, + proposal=proposal, + req_num=2, + accept_lengths_cpu=torch.tensor([1, 1], dtype=torch.int32), + ) first_plan = planner.plan(req_num=2, original_batch_size=8) - planner.update_predicted_schedule_probs(schedule_probs=schedule_probs, req_num=2) + engine.update_planner_feedback( + plan=plan, + proposal=proposal, + req_num=2, + accept_lengths_cpu=torch.tensor([1, 1], dtype=torch.int32), + ) second_plan = planner.plan(req_num=2, original_batch_size=8) - assert first_plan.dynamic_batch_size == 2 + assert first_plan.dynamic_batch_size == 8 assert second_plan.dynamic_batch_size == 4 - assert second_plan.draft_step == 3 + assert second_plan.draft_step == second_plan.pre_draft_step == 3 + + +def test_dspark_uses_confidence_scores_without_acceptance_ema(): + planner = build_dspark_planner() + for batch_size, target_cost in ((2, 1.0), (4, 1.2), (8, 10.0)): + planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) + planner.update_infer_cost(batch_size=6, infer_cost_ms=0.5, is_draft_model=True) + + low_confidence = np.cumprod(np.full((2, 3), 0.1), axis=1) + high_confidence = np.cumprod(np.full((2, 3), 0.9), axis=1) + + assert planner._select_dynamic_batch_size_from_survival_scores(2, low_confidence) == 2 + assert planner._select_dynamic_batch_size_from_survival_scores(2, high_confidence) == 4 def test_topk_prefix_sums_only_computes_requested_counts(): - result = DSparkDynamicSpecPlanner._topk_prefix_sums( + result = DSparkPlanner._topk_prefix_sums( values=np.asarray([0.1, 0.9, 0.4, 0.7]), counts=[0, 2, 4], ) diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 0306487c10..c85c355bcd 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -163,7 +163,7 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): draft_models=[draft_model], _gen_argmax_token_ids_and_prob=lambda _: (flat_token_ids, flat_token_probs), ) - proposer = DFlashProposer(SimpleNamespace(backend=backend, enable_dynamic_spec=True)) + proposer = DFlashProposer(backend=backend, enable_dynamic_spec=True) proposer.extend_draft_kv_cache = lambda **_: None proposer.select_accepted_tail_rows = lambda **_: selected_rows proposer.build_block_draft_input = lambda **_: (SimpleNamespace(), torch.tensor([10, 11])) @@ -177,13 +177,14 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): accept_len=torch.tensor([1, 1]), ) - expected_blocks = flat_token_ids.reshape(2, block_size)[:, 1:3] + expected_blocks = flat_token_ids.reshape(2, block_size)[:, :2] torch.testing.assert_close(proposal.token_ids[selected_rows, 1:], expected_blocks) - assert len(proposal.draft_probs) == 2 - for step, probs in enumerate(proposal.draft_probs): - expected_probs = torch.zeros(verify_row_count) - expected_probs[selected_rows] = flat_token_probs.reshape(2, block_size)[:, step + 1] - torch.testing.assert_close(probs, expected_probs) + assert proposal.schedule_scores.shape == (verify_row_count, 2) + for step in range(2): + scores = proposal.schedule_scores[:, step] + expected_scores = torch.zeros(verify_row_count) + expected_scores[selected_rows] = flat_token_probs.reshape(2, block_size)[:, step] + torch.testing.assert_close(scores, expected_scores) def test_hidden_collector_reads_target_layer_ids(monkeypatch): From 2870d32e271b933f512b0f92cab5087d7aecd56d Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Fri, 14 Aug 2026 06:28:14 +0000 Subject: [PATCH 010/103] fix: stabilize dynamic speculative scheduling --- .../common/basemodel/attention/base_att.py | 2 +- lightllm/common/basemodel/batch_objs.py | 4 +- lightllm/common/basemodel/infer_struct.py | 2 +- .../basemodel/triton_kernel/mtp_utils.py | 5 +- lightllm/models/qwen2_vl/infer_struct.py | 6 +- lightllm/models/qwen3_5_dflash/model.py | 4 +- lightllm/models/qwen3_5_dspark/model.py | 7 +- .../layer_infer/pre_layer_infer.py | 2 +- lightllm/server/api_cli.py | 3 +- .../mode_backend/dp_backend/impl.py | 10 +- .../mode_backend/mtp_pre_process.py | 22 +- .../router/model_infer/speculative/planner.py | 107 ++-- .../model_infer/speculative/proposers/base.py | 42 +- .../speculative/proposers/dflash.py | 209 +------ .../speculative/proposers/dspark.py | 70 +-- .../speculative/proposers/eagle3.py | 202 +----- .../speculative/proposers/eagle_mtp.py | 578 ++++++++---------- .../speculative/proposers/parallel_block.py | 118 ++++ .../speculative/proposers/vanilla_mtp.py | 94 ++- .../common/basemodel/test_hidden_collector.py | 8 +- .../models/qwen2_vl/test_infer_struct.py | 29 + .../models/test_qwen3_dspark_model_output.py | 2 +- .../speculative/test_eagle_overlap.py | 58 +- .../model_infer/speculative/test_planner.py | 90 ++- unit_tests/utils/test_speculative_utils.py | 82 ++- 25 files changed, 803 insertions(+), 953 deletions(-) create mode 100644 lightllm/server/router/model_infer/speculative/proposers/parallel_block.py create mode 100644 unit_tests/models/qwen2_vl/test_infer_struct.py diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index da897faef0..cad0563f20 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -45,7 +45,7 @@ def uses_dynamic_spec_verify_layout(self, infer_state: "InferStateInfo") -> bool return False # Target verification may compact each request to a different row count. - # Block draft forwards still use their checkpoint-defined fixed layout. + # Parallel block drafter forwards still use their checkpoint-defined fixed layout. return args.mtp_mode not in ("dspark", "dflash") or not self.model.is_mtp_draft_model def _find_layer_index( diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index abf0e6cdd1..69d3ef369c 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -37,8 +37,7 @@ class ModelInput: mem_indexes: torch.Tensor = None is_prefill: bool = False b_ready_cache_len: torch.Tensor = None - # 只会在继承 Qwen2VLInferStateInfo 的 MRoPE 模型 decode 阶段使用,如 - # Qwen2/2.5-VL、Qwen3-VL/MOE/Omni、Qwen3.5;普通模型不会使用。 + # MRoPE position offset; preserve it across row-aligned input transforms. b_position_delta: torch.Tensor = None b_prefill_start_loc: torch.Tensor = None multimodal_params: list = None @@ -75,7 +74,6 @@ def to_cuda(self): self.b_ready_cache_len = self.b_ready_cache_len.cuda(non_blocking=True) if self.b_position_delta is not None: self.b_position_delta = self.b_position_delta.cuda(non_blocking=True) - assert self.is_prefill is False, "b_position_delta should only be used in decode phase." else: assert self.is_prefill is True, "decode ModelInput should provide b_position_delta." diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index ac9922ae69..28d8a5f099 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -38,7 +38,7 @@ def __init__(self): self.b_mark_shared_group: torch.Tensor = None # only for diverse mode used in decode phase. self.b_mtp_index: torch.Tensor = None - # only for mrope model in decode phase used. + # MRoPE position offset propagated from ModelInput. self.b_position_delta: torch.Tensor = None self.b_seq_len: torch.Tensor = None diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index 8cee6e5163..cb8a3dc027 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -528,8 +528,9 @@ def prepare_dynamic_spec_model_input( # Keep CPU mem_indexes unfiltered here. Copying selected_row_mask back to # CPU in this hot path synchronizes the overlap stream; the router frees # unselected/rejected CPU mem indexes after its existing async mask copy is - # consumed. Decode only needs b_position_delta on device, so placeholder - # multimodal metadata keeps padded ModelInput checks shape-consistent. + # consumed. Decode and draft-cache commit use the compacted + # b_position_delta, so placeholder multimodal metadata only needs to keep + # ModelInput shapes consistent. if model_input.multimodal_params is not None: # Read-only placeholders: avoid rebuilding hundreds of nested Python # objects on every compacted decode iteration. diff --git a/lightllm/models/qwen2_vl/infer_struct.py b/lightllm/models/qwen2_vl/infer_struct.py index 04f7bc3895..f3ae4ba668 100644 --- a/lightllm/models/qwen2_vl/infer_struct.py +++ b/lightllm/models/qwen2_vl/infer_struct.py @@ -17,10 +17,14 @@ def init_some_extra_state(self, model): rope_scaling = model.config.get("rope_scaling", {}) self.rope_type = rope_scaling.get("rope_type", rope_scaling.get("type", None)) InferStateInfo.init_some_extra_state(self, model) - if self.is_prefill: + if self.is_prefill and self.b_position_delta is None: self.position_ids = self.get_mrope_position(self.multimodal_params) else: b_position_delta = self.b_position_delta.to(dtype=self.position_ids.dtype) + if self.is_prefill: + # Speculative draft commits use one prefill row per verified + # token, whose decode-time MRoPE offset is already available. + assert b_position_delta.shape == self.position_ids.shape position_ids = self.position_ids + b_position_delta self.position_ids = position_ids.unsqueeze(0).expand(3, -1) diff --git a/lightllm/models/qwen3_5_dflash/model.py b/lightllm/models/qwen3_5_dflash/model.py index a9f3ce71cd..da3755588d 100644 --- a/lightllm/models/qwen3_5_dflash/model.py +++ b/lightllm/models/qwen3_5_dflash/model.py @@ -8,7 +8,7 @@ @DraftModelRegistry(model_type=("qwen3_5", "qwen3_5_text"), spec_modes="dflash") class Qwen3_5DFlashModel(Qwen3DFlashModel): - """Qwen3.5 DFlash draft model.""" + """Adapter for a Qwen3 DFlash checkpoint paired with a Qwen3.5 target.""" pre_and_post_weight_class = Qwen35DFlashPreAndPostLayerWeight @@ -38,7 +38,7 @@ def _init_mem_manager(self): target_mem_manager.head_dim, ) assert draft_kv_shape == target_kv_shape, ( - "Qwen3.5 block draft requires matching draft and target KV shapes, " + "Qwen3.5 parallel block drafter requires matching draft and target KV shapes, " f"got draft={draft_kv_shape}, target={target_kv_shape}." ) super()._init_mem_manager() diff --git a/lightllm/models/qwen3_5_dspark/model.py b/lightllm/models/qwen3_5_dspark/model.py index cec8b7b8c9..23182ce0f6 100644 --- a/lightllm/models/qwen3_5_dspark/model.py +++ b/lightllm/models/qwen3_5_dspark/model.py @@ -5,7 +5,7 @@ @DraftModelRegistry(model_type=("qwen3_5", "qwen3_5_text"), spec_modes="dspark") class Qwen3_5DSparkModel(Qwen3DSparkModel): - """Qwen3 DSpark draft model paired with a Qwen3.5 target.""" + """Adapter for the current Qwen3 DSpark checkpoint with a Qwen3.5 target.""" def _init_config(self): super()._init_config() @@ -15,8 +15,7 @@ def _init_config(self): if "rope_theta" in rope_parameters and "rope_theta" not in self.config: self.config["rope_theta"] = rope_parameters["rope_theta"] - # DeepSpec trains this Qwen3 draft with ordinary 1D full-head RoPE. Do - # not pass the Qwen3.5 target's MRoPE layout into the draft backbone. + # Match the draft checkpoint's 1D full-head RoPE layout. self.config["rope_scaling"] = None self.config["partial_rotary_factor"] = 1.0 @@ -34,7 +33,7 @@ def _init_mem_manager(self): target_mem_manager.head_dim, ) assert draft_kv_shape == target_kv_shape, ( - "Qwen3.5 block draft requires matching draft and target KV shapes, " + "Qwen3.5 parallel block drafter requires matching draft and target KV shapes, " f"got draft={draft_kv_shape}, target={target_kv_shape}." ) super()._init_mem_manager() diff --git a/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py index b52a61f945..f437a7d28c 100644 --- a/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py +++ b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py @@ -17,7 +17,7 @@ def prepare_spec_draft_hiddens( ) -> None: target_hiddens = infer_state.mtp_draft_input_hiddens # Target verification provides concatenated auxiliary-layer hiddens (N * H). - # Recurrent draft steps feed the previous draft output, which is already H. + # Autoregressive draft steps feed the previous draft output, which is already H. if target_hiddens.shape[-1] != self.hidden_size_: target_hiddens = layer_weight.fc_weight_.mm(target_hiddens) infer_state.eagle_draft_hidden_states = target_hiddens diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 07f8ace1de..ffc5434e28 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -745,7 +745,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: default=None, help="""Speculative decoding mode. *_with_att and *_no_att select attention or non-attention draft models; - eagle3 uses a recurrent EAGLE3 draft; dspark and dflash use block draft models.""", + eagle3 uses autoregressive EAGLE-3 drafting; dflash uses block-diffusion drafting; + dspark uses semi-autoregressive parallel drafting.""", ) parser.add_argument( "--mtp_draft_model_dir", diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 335ab04a48..04ab0de46d 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -29,8 +29,8 @@ def __init__(self) -> None: spec_mode = get_env_start_args().mtp_mode if spec_mode is not None: if spec_mode in ("dspark", "dflash"): - raise NotImplementedError("DP backend does not support DFlash/DSpark block draft mode yet.") - self.uses_recurrent_draft = spec_mode in ( + raise NotImplementedError("DP backend does not support DFlash/DSpark parallel block drafting yet.") + self.uses_autoregressive_drafter = spec_mode in ( "eagle_with_att", "eagle_no_att", "eagle3", @@ -43,13 +43,13 @@ def __init__(self) -> None: self.decode = self.decode_overlap_mtp self._draft_decode_overlap_func = ( self._draft_decode_eagle_overlap - if self.uses_recurrent_draft + if self.uses_autoregressive_drafter else self._draft_decode_vanilla_overlap ) else: self.decode = self.decode_mtp self._draft_decode_func = ( - self._draft_decode_eagle if self.uses_recurrent_draft else self._draft_decode_vanilla + self._draft_decode_eagle if self.uses_autoregressive_drafter else self._draft_decode_vanilla ) else: if self.enable_prefill_microbatch_overlap: @@ -680,7 +680,7 @@ def _draft_decode_eagle( # DP keeps the target verify layout padded for collective shape # agreement. The proposer still follows the common topology: one - # full-row extend, followed by recurrent decode over one row per + # full-row extend, followed by autoregressive drafting over one row per # (real or HOLD) request. proposal = self.spec_engine.propose_next( main_model_input=model_input, diff --git a/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py b/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py index 3ef0395431..2e392e6ffa 100644 --- a/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py @@ -1,24 +1,20 @@ import torch -import copy from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids def prepare_mtp_prefill_inputs( - model_input: ModelInput, b_next_token_ids: torch.Tensor, mtp_draft_input_hiddens: torch.Tensor -): - # enable_prefill_decode_mixed 模式下,decode 请求混合在 prefill 请求中。 - # 但是mtp的input_ids已经是恢复ok,已经是正常的input_ids, 所以移除掉 b_is_decode_req。 - # 防止在 forward 阶段,因为 b_is_decode_req 不为空,导致 input_ids 被特殊处理。 - new_model_input = copy.copy(model_input) - new_model_input.b_is_decode_req = None - - new_input_ids = gen_mtp_new_input_ids( + model_input: ModelInput, + b_next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, +) -> ModelInput: + # MTP supplies explicit token ids; mixed-prefill gathering must not replace them. + model_input.b_is_decode_req = None + model_input.input_ids = gen_mtp_new_input_ids( input_ids=model_input.input_ids, b_next_token_ids=b_next_token_ids, b_seq_len=model_input.b_seq_len, b_ready_cache_len=model_input.b_ready_cache_len, ) - new_model_input.input_ids = new_input_ids - new_model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens - return new_model_input + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + return model_input diff --git a/lightllm/server/router/model_infer/speculative/planner.py b/lightllm/server/router/model_infer/speculative/planner.py index d8c954caa6..5b24701c83 100644 --- a/lightllm/server/router/model_infer/speculative/planner.py +++ b/lightllm/server/router/model_infer/speculative/planner.py @@ -65,12 +65,11 @@ class LightSpecPlanner: LightSpec evaluates a runtime configuration as ``(N, B, d)``: logical requests, physical target rows, and draft configuration. Planning follows - the paper's budget-then-fill order. It first chooses ``B`` for the proposal - produced by ``pre_draft_step``, then keeps that ``B`` fixed while choosing - the ``draft_step`` that will produce the next proposal. Candidate identities - and per-request verify widths belong to the GPU Fill stage, not this planner. - The next draft depth bounds the next iteration's verify budget; it does not - need to reproduce the verify budget consumed in the current iteration. + the paper's budget-then-fill order. It chooses the current ``B`` within the + proposal produced by ``pre_draft_step``. The next ``draft_step`` is selected + from self-consistent ``(B, d)`` configurations for the following iteration, + so a short current proposal cannot prevent the planner from drafting deeper + again. Candidate identities belong to the GPU Fill stage. Each configuration minimizes estimated milliseconds per committed token: ``(target_cost(B) + draft_cost(N, B, d)) / expected_progress(N, B, d)``. @@ -99,6 +98,10 @@ def __init__( # complete batch. The draft configuration is part of the key because a # deeper candidate pool may improve Fill even at the same (N, B). self.progress_ema_by_config: Dict[Tuple[int, int, int], _EMAValue] = {} + # Full-width verification provides an unbiased survival probability + # for every draft depth. It is the fallback for configurations that + # have not produced their own progress observation yet. + self.prefix_survival_by_depth: List[Optional[_EMAValue]] = [None] * self.max_draft_step # The current verify width is bounded by the proposal built last time. self.pre_draft_step = self.max_draft_step @@ -147,10 +150,7 @@ def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> for dynamic_batch_size in batch_sizes ] dynamic_batch_size = batch_sizes[np.argmin(costs)] - draft_step = self._select_draft_step( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - ) + draft_step = self._select_draft_step(req_num=req_num) self.pre_draft_step = draft_step return SpecDecodePlan( @@ -202,25 +202,44 @@ def update_verified_batch( config = (int(req_num), int(dynamic_batch_size), int(verified_draft_step)) progress = float(accept_lengths.sum()) / dynamic_batch_size if config not in self.progress_ema_by_config: - init_progress = progress if not self.progress_ema_by_config else self._estimate_nearby_progress(*config) self.progress_ema_by_config[config] = _EMAValue( decay=0.9, - init_value=init_progress, + init_value=progress, ) self.progress_ema_by_config[config].update(progress) - def _select_draft_step(self, req_num: int, dynamic_batch_size: int) -> int: + if dynamic_batch_size == req_num * (verified_draft_step + 1): + for depth in range(1, verified_draft_step + 1): + survival = float(np.mean(accept_lengths > depth)) + survival_ema = self.prefix_survival_by_depth[depth - 1] + if survival_ema is None: + survival_ema = _EMAValue(decay=0.9, init_value=survival) + self.prefix_survival_by_depth[depth - 1] = survival_ema + survival_ema.update(survival) + + def _select_draft_step(self, req_num: int) -> int: + """Choose the best self-consistent next proposal depth. + + The current target batch is bounded by ``pre_draft_step``, but the + proposal generated now is consumed by the next iteration. Evaluate + each candidate depth with the verify widths that depth can create so a + short current proposal cannot permanently prevent deeper drafting. + """ + best_cost_ms = float("inf") best_draft_step = self.draft_steps[0] for draft_step in self.draft_steps: - cost_ms = self._get_cost_ms( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - ) - if cost_ms < best_cost_ms: - best_cost_ms = cost_ms - best_draft_step = draft_step + max_batch_size = req_num * (draft_step + 1) + batch_sizes = self.target_infer_costs.get_batch_size_keys_between(req_num, max_batch_size) + for dynamic_batch_size in batch_sizes: + cost_ms = self._get_cost_ms( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + ) + if cost_ms < best_cost_ms: + best_cost_ms = cost_ms + best_draft_step = draft_step return best_draft_step def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: @@ -245,36 +264,23 @@ def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) - def _estimate_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: config = (int(req_num), int(dynamic_batch_size), int(draft_step)) ema = self.progress_ema_by_config.get(config) - return ema.get() if ema is not None else self._estimate_nearby_progress(*config) + return ema.get() if ema is not None else self._estimate_prefix_progress(*config) - def _estimate_nearby_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: - if not self.progress_ema_by_config: + def _estimate_prefix_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: + if not any(self.prefix_survival_by_depth): return 1.0 - rows_per_req = dynamic_batch_size / req_num - - def distance(config: Tuple[int, int, int]): - observed_req_num, observed_batch_size, observed_draft_step = config - width_distance = (observed_batch_size / observed_req_num - rows_per_req) / (self.max_draft_step + 1) - draft_distance = (observed_draft_step - draft_step) / max(self.max_draft_step, 1) - request_distance = (observed_req_num - req_num) / max(observed_req_num, req_num) - return width_distance ** 2 + draft_distance ** 2 + request_distance ** 2 - - nearest_config = min(self.progress_ema_by_config, key=distance) - observed_req_num, observed_batch_size, _ = nearest_config - observed_rows_per_req = observed_batch_size / observed_req_num - observed_progress_per_req = observed_rows_per_req * self.progress_ema_by_config[nearest_config].get() - - # Fill keeps the highest-survival rows when B shrinks. Transfer the - # neighbor's committed progress per request, then normalize it for the - # target B. Copying rho directly would predict that useful progress - # falls in proportion to B and make B=N an artificial absorbing state. - estimated_progress_per_req = min( - max(observed_progress_per_req, 1.0), - rows_per_req, - draft_step + 1, - ) - return estimated_progress_per_req / rows_per_req + remaining_draft_rows = dynamic_batch_size - req_num + expected_tokens = float(req_num) + for depth in range(min(draft_step, self.max_draft_step)): + selected_rows = min(req_num, remaining_draft_rows) + if selected_rows <= 0: + break + survival_ema = self.prefix_survival_by_depth[depth] + survival = 1.0 if survival_ema is None else survival_ema.get() + expected_tokens += selected_rows * survival + remaining_draft_rows -= selected_rows + return expected_tokens / dynamic_batch_size class DSparkPlanner: @@ -474,8 +480,9 @@ def estimate(self, batch_size: int) -> float: def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> List[int]: start = min(int(batch_size1), int(batch_size2)) end = max(int(batch_size1), int(batch_size2)) - batch_sizes = list(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) - return batch_sizes or [end] + batch_sizes = set(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) + batch_sizes.update((start, end)) + return sorted(batch_sizes) class _EMAValue: diff --git a/lightllm/server/router/model_infer/speculative/proposers/base.py b/lightllm/server/router/model_infer/speculative/proposers/base.py index 6bcc1be004..b61847ee01 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/speculative/proposers/base.py @@ -61,23 +61,6 @@ def get_draft_cost_ms( raise NotImplementedError - def select_accepted_tail_rows(self, b_req_mtp_start_loc: torch.Tensor, accept_len: torch.Tensor) -> torch.Tensor: - return (b_req_mtp_start_loc + accept_len - 1).to(torch.long) - - def scatter_selected_step_probs( - self, - selected_rows: torch.Tensor, - selected_probs: torch.Tensor, - verify_row_count: int, - ) -> torch.Tensor: - out = torch.zeros( - (verify_row_count, *selected_probs.shape[1:]), - dtype=torch.float32, - device=selected_probs.device, - ) - out[selected_rows] = selected_probs.float() - return out - def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: """Allocate draft-owned temporary KV slots.""" @@ -106,10 +89,8 @@ def build_draft_state_from_prefill( by the selected speculative algorithm. - `next_token_ids`: first accepted target token, shape [run_req_num]. - This hook only prepares draft-side state. It intentionally does not - scatter proposal tokens; the first decode iteration verifies as having - no draft candidates and produces the first proposal through - `propose_next`. + This hook only prepares draft-side state. It does not create proposal + tokens; the first decode iteration creates them through `propose_next`. """ raise NotImplementedError @@ -138,20 +119,11 @@ def propose_next( ) -> SpecProposal: """Generate candidate tokens after one target decode forward. - Inputs: - - `main_model_input`: target decode ModelInput. In the fixed layout its - batch is laid out as [req0-main, req0-draft1, ...]. Dynamic scheduling may - compact this batch before target forward. - - `next_token_ids`: target sampled ids for rows in `main_model_input`, - shape [verify_batch]. - - `b_req_mtp_start_loc`: start row for each logical request inside the - speculative verify batch, shape [logical_req_num]. - - `draft_step`: number of candidate draft tokens to produce. - - `accept_len`: optional accepted-prefix length from the just-finished - target forward. Stateful block proposers use it to commit the - accepted target-hidden segment before preparing the next block. - - Returns a SpecProposal whose `token_ids[:, 0]` is `next_token_ids`. + `main_model_input` contains the target verify rows, possibly compacted + by dynamic scheduling. `b_req_mtp_start_loc` identifies each logical + request's first row. + + Column 0 of the returned proposal must equal `next_token_ids`. """ raise NotImplementedError diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index a67633d4f2..c700033e21 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -1,61 +1,19 @@ from __future__ import annotations -import copy - import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.speculative.proposers.parallel_block import ParallelBlockProposer -class DFlashProposer(BaseSpecProposer): - """Non-causal block proposer for DFlash. +class DFlashProposer(ParallelBlockProposer): + """DFlash block-diffusion proposer. - DFlash remains a non-causal block-prefill draft model, not a recurrent - token decoder. The service flow is: - - verify target tokens - - extend the DFlash draft KV cache with target hidden rows - - draft a new non-causal block from the accepted tail row - The memory lifecycle stays in the normal decode free path: - - rejected target token slots are freed by the normal decode free path - - current-block scratch KV uses extra mem slots returned through - `SpecProposal.extra_mem_indexes_cpu` + The drafter predicts a complete token block in one parallel forward from + the accepted-tail anchor and mask-token positions. """ - def get_draft_steps(self): - return (self.backend.max_draft_step,) - - def get_draft_cost_ms( - self, - draft_infer_costs, - req_num: int, - verify_batch_size: int, - draft_step: int, - ) -> float: - block_size = self.backend.draft_models[0].block_size - # Block drafting first commits all verified rows, then generates one - # complete checkpoint-defined block for each request. - extend_cost_ms = draft_infer_costs.estimate(verify_batch_size) - block_cost_ms = draft_infer_costs.get(req_num * block_size) - return extend_cost_ms + block_cost_ms - - @torch.no_grad() - def build_draft_state_from_prefill( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - next_token_ids: torch.Tensor, - ) -> None: - target_hidden = target_model_output.spec_hidden - if target_hidden.numel() == 0: - return - - draft_model = self.backend.draft_models[0] - draft_input = copy.copy(target_model_input) - # DFlash consumes target hidden states directly on this prefill path. - draft_input.mtp_draft_input_hiddens = target_hidden - draft_model.forward(draft_input) - @torch.no_grad() def propose_next( self, @@ -66,151 +24,48 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> SpecProposal: - num_reqs = int(b_req_mtp_start_loc.shape[0]) + request_count = int(b_req_mtp_start_loc.shape[0]) + verify_row_count = int(next_token_ids.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - token_ids = next_token_ids.new_full( - (next_token_ids.shape[0], draft_step + 1), + proposal_token_ids = next_token_ids.new_full( + (verify_row_count, draft_step + 1), fill_value=1, ) - token_ids[:, 0] = next_token_ids + proposal_token_ids[:, 0] = next_token_ids - self.extend_draft_kv_cache( + # One accepted-tail anchor expands to a complete block-diffusion draft. + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + draft_input, extra_mem_indexes_cpu = self.build_block_draft_input( main_model_input=main_model_input, - target_hidden=main_model_output.spec_hidden, - ) - - if draft_step == 0: - return SpecProposal( - token_ids=token_ids, - extra_mem_indexes_cpu=None, - ) - - # DFlash drafts from the accepted tail row of each request; one anchor - # row expands to a complete non-causal draft block. - selected_rows = self.select_accepted_tail_rows( - b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=accept_len, + next_token_ids=next_token_ids, + accepted_tail_rows=accepted_tail_rows, + request_count=request_count, ) - draft_input, draft_mem_indexes_cpu = self.build_block_draft_input( + self.extend_draft_kv_cache( main_model_input=main_model_input, - next_token_ids=next_token_ids, - selected_rows=selected_rows, - num_reqs=num_reqs, + target_hidden=main_model_output.spec_hidden, ) - draft_model_output = draft_model.forward(draft_input) + draft_output = draft_model.forward(draft_input) if self.enable_dynamic_spec: - flat_token_ids, flat_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) + flat_draft_token_ids, flat_draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) else: - flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) - block_token_ids = flat_token_ids.reshape(num_reqs, block_size) - token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] + flat_draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) + proposal_token_ids[accepted_tail_rows, 1:] = block_draft_token_ids[:, :draft_step] schedule_scores = None if self.enable_dynamic_spec: - block_token_probs = flat_token_probs.reshape(num_reqs, block_size) - selected_token_probs = block_token_probs[:, :draft_step] - schedule_scores = self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=selected_token_probs, - verify_row_count=next_token_ids.shape[0], + block_draft_token_probs = flat_draft_token_probs.reshape(request_count, block_size) + schedule_scores = torch.zeros( + (verify_row_count, draft_step), + dtype=torch.float32, + device=next_token_ids.device, ) + schedule_scores[accepted_tail_rows] = block_draft_token_probs[:, :draft_step].float() return SpecProposal( - token_ids=token_ids, - extra_mem_indexes_cpu=draft_mem_indexes_cpu, + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, schedule_scores=schedule_scores, ) - - def extend_draft_kv_cache(self, main_model_input: ModelInput, target_hidden: torch.Tensor) -> None: - draft_model = self.backend.draft_models[0] - batch_size = int(target_hidden.shape[0]) - - draft_kv_input = copy.copy(main_model_input) - draft_kv_input.batch_size = batch_size - draft_kv_input.total_token_num = batch_size - empty_multimodal_params = {"images": [], "audios": []} - draft_kv_input.multimodal_params = [empty_multimodal_params] * batch_size - # This hidden-commit prefill path does not consume token ids, but - # InferState uses input_ids.shape[0] to build position ids. Keep it - # aligned with the fixed-shape target verify batch. - draft_kv_input.input_ids = torch.empty( - (batch_size,), - dtype=torch.int64, - device=target_hidden.device, - ) - draft_kv_input.max_q_seq_len = 1 - draft_kv_input.prefix_total_token_num = 0 - draft_kv_input.is_prefill = True - # Match Eagle3's fixed verify commit: write every speculative row and - # let the accepted-tail sequence length select the valid prefix. Rejected - # suffix slots are ignored and released by the normal verify free path. - # Keeping the batch shape static avoids torch.nonzero's implicit D2H - # synchronization and lets the host enqueue the draft work immediately. - draft_kv_input.b_ready_cache_len = main_model_input.b_seq_len - 1 - draft_kv_input.b_prefill_start_loc = torch.arange( - batch_size, - dtype=torch.int32, - device=target_hidden.device, - ) - draft_kv_input.b_position_delta = None - draft_kv_input.mtp_draft_input_hiddens = target_hidden - draft_model.forward(draft_kv_input) - - def build_block_draft_input( - self, - main_model_input: ModelInput, - next_token_ids: torch.Tensor, - selected_rows: torch.Tensor, - num_reqs: int, - ): - draft_model = self.backend.draft_models[0] - block_size = int(draft_model.block_size) - draft_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * block_size) - - draft_input_ids = next_token_ids.new_full( - (num_reqs * block_size,), - fill_value=draft_model.mask_token_id, - ) - # Each block is [accepted_token, mask, ..., mask], matching DeepSpec's - # draft input. The proposer later maps the block logits back to - # [base_token + draft_tokens] for target verification. - draft_input_ids[::block_size] = next_token_ids.index_select(0, selected_rows) - - block_offsets = torch.arange( - block_size, - dtype=main_model_input.b_seq_len.dtype, - device=next_token_ids.device, - ) - draft_input = copy.copy(main_model_input) - draft_input.input_ids = draft_input_ids - draft_input.total_token_num = draft_input.input_ids.shape[0] - draft_input.batch_size = draft_input.total_token_num - draft_input.max_q_seq_len = 1 - draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size - draft_input.draft_step = block_size - 1 - draft_input.b_req_idx = ( - main_model_input.b_req_idx.index_select(0, selected_rows).repeat_interleave(block_size).contiguous() - ) - draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) - # b_seq_len is real metadata, not cosmetic: copy_kv_index_to_req and - # FA3 use it to place scratch KV and compute the block cache length. - draft_input.b_seq_len = ( - (main_model_input.b_seq_len.index_select(0, selected_rows)[:, None] + block_offsets[None, :] + 1) - .reshape(-1) - .contiguous() - ) - if main_model_input.b_position_delta is not None: - draft_input.b_position_delta = ( - main_model_input.b_position_delta.index_select(0, selected_rows) - .repeat_interleave(block_size) - .contiguous() - ) - else: - draft_input.b_position_delta = torch.zeros_like(draft_input.b_req_idx) - draft_input.mem_indexes = draft_mem_indexes_cpu.cuda(non_blocking=True) - draft_input.b_mark_shared_group = torch.zeros_like(draft_input.b_req_idx) - draft_input.b_mark_shared_group[block_size - 1 :: block_size] = block_size - empty_multimodal_params = {"images": [], "audios": []} - draft_input.multimodal_params = [empty_multimodal_params] * draft_input.batch_size - return draft_input, draft_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/speculative/proposers/dspark.py index fc40151ed4..641cade2d0 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dspark.py @@ -5,15 +5,15 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer +from lightllm.server.router.model_infer.speculative.proposers.parallel_block import ParallelBlockProposer -class DSparkProposer(DFlashProposer): - """DSpark block proposer. +class DSparkProposer(ParallelBlockProposer): + """DSpark semi-autoregressive parallel-block proposer. - DSpark shares DFlash's target-hidden KV injection and non-causal block - backbone. Its post layer returns Markov-corrected logits and optional - confidence logits, so the proposer follows the same token path as DFlash. + A parallel DFlash-style backbone generates the block features in one pass; + a lightweight sequential Markov head adds intra-block token dependency. + The confidence head supplies per-position scheduling scores. """ @torch.no_grad() @@ -26,8 +26,8 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> SpecProposal: - num_reqs = int(b_req_mtp_start_loc.shape[0]) - verify_row_count = next_token_ids.shape[0] + request_count = int(b_req_mtp_start_loc.shape[0]) + verify_row_count = int(next_token_ids.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) proposal_token_ids = next_token_ids.new_full( @@ -36,7 +36,7 @@ def propose_next( ) proposal_token_ids[:, 0] = next_token_ids schedule_scores = ( - torch.empty( + torch.zeros( (verify_row_count, draft_step), dtype=torch.float32, device=next_token_ids.device, @@ -45,47 +45,39 @@ def propose_next( else None ) - self.extend_draft_kv_cache( + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + draft_input, extra_mem_indexes_cpu = self.build_block_draft_input( main_model_input=main_model_input, - target_hidden=main_model_output.spec_hidden, - ) - - if draft_step == 0: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - schedule_scores=schedule_scores, - ) - - selected_rows = self.select_accepted_tail_rows( - b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=accept_len, + next_token_ids=next_token_ids, + accepted_tail_rows=accepted_tail_rows, + request_count=request_count, ) - draft_input, draft_mem_indexes_cpu = self.build_block_draft_input( + self.extend_draft_kv_cache( main_model_input=main_model_input, - next_token_ids=next_token_ids, - selected_rows=selected_rows, - num_reqs=num_reqs, + target_hidden=main_model_output.spec_hidden, ) - draft_model_output = draft_model.forward(draft_input) + draft_output = draft_model.forward(draft_input) - if draft_model_output.draft_token_ids is None: - flat_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) + if draft_output.draft_token_ids is None: + flat_draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) else: - flat_token_ids = draft_model_output.draft_token_ids - block_token_ids = flat_token_ids.reshape(num_reqs, block_size) - proposal_token_ids[selected_rows, 1:] = block_token_ids[:, :draft_step] + flat_draft_token_ids = draft_output.draft_token_ids + block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) + proposal_token_ids[accepted_tail_rows, 1:] = block_draft_token_ids[:, :draft_step] if self.enable_dynamic_spec: - confidence_logits = draft_model_output.confidence_logits + confidence_logits = draft_output.confidence_logits if confidence_logits is None: raise RuntimeError("DSpark dynamic verify requires confidence head logits") # Match the clamp used by the GPU dynamic row selector before it # converts conditional confidence to prefix survival probability. - schedule_scores = self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=confidence_logits[:, :draft_step].sigmoid().clamp(min=0.01, max=0.99), - verify_row_count=verify_row_count, + schedule_scores[accepted_tail_rows] = ( + confidence_logits[:, :draft_step] + .sigmoid() + .clamp( + min=0.01, + max=0.99, + ) ) schedule_scores_cpu = None @@ -97,7 +89,7 @@ def propose_next( return SpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=draft_mem_indexes_cpu, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, schedule_scores=schedule_scores, schedule_scores_cpu=schedule_scores_cpu, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py index c150f56eb4..453eaa5738 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle3.py @@ -1,210 +1,16 @@ from __future__ import annotations -import math - import torch -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import RecurrentEagleMTPProposer +from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import AutoregressiveEagleProposer -class Eagle3Proposer(RecurrentEagleMTPProposer): +class Eagle3Proposer(AutoregressiveEagleProposer): """Eagle3 proposer. - After target verification, Eagle3 commits the accepted target segment into - the draft cache with target hidden states, then drafts the next proposal - from that corrected state. + It uses the shared autoregressive proposal flow with an additional + draft-to-target vocabulary mapping step. """ - _DRAFT_PRUNE_SAFETY_FACTOR = 1.10 - _DRAFT_PRUNE_MIN_DEPTH = 4 - - def get_draft_cost_ms( - self, - draft_infer_costs, - req_num: int, - verify_batch_size: int, - draft_step: int, - ) -> float: - draft_cost_ms = draft_infer_costs.estimate(verify_batch_size) - active_count = req_num - draft_row_budget = max(1, verify_batch_size - req_num) - for step in range(1, draft_step): - active_count = self._get_pruned_active_count( - current_count=active_count, - draft_row_budget=draft_row_budget, - next_depth=step + 1, - ) - draft_cost_ms += draft_infer_costs.get(active_count) - return draft_cost_ms - - def _get_pruned_active_count( - self, - current_count: int, - draft_row_budget: int, - next_depth: int, - ) -> int: - if next_depth < self._DRAFT_PRUNE_MIN_DEPTH or current_count <= 1: - return current_count - - # A selected token at depth d consumes all d prefix draft rows from - # that request. Therefore at most L/d chains can reach depth d when - # the next target verify has L draft-row slots. Keep a small safety - # margin, then retain the highest-survival chain frontiers. - active_count = math.ceil(self._DRAFT_PRUNE_SAFETY_FACTOR * max(1, draft_row_budget) / next_depth) - return min(current_count, max(1, active_count)) - def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) - - def propose_next( - self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> SpecProposal: - verify_row_count = next_token_ids.shape[0] - num_reqs = b_req_mtp_start_loc.shape[0] - proposal_token_ids = next_token_ids.new_full( - (verify_row_count, draft_step + 1), - fill_value=1, - ) - proposal_token_ids[:, 0].copy_(next_token_ids) - collect_dynamic_probs = self.enable_dynamic_spec - schedule_scores = ( - torch.zeros( - (verify_row_count, draft_step), - dtype=torch.float32, - device=next_token_ids.device, - ) - if collect_dynamic_probs - else None - ) - - target_hidden = main_model_output.spec_hidden - - # Scatter consumes the accepted-tail row for each request; only those - # rows need new draft columns after the commit step. - selected_rows = self.select_accepted_tail_rows( - b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=accept_len, - ) - draft_model = self.backend.draft_models[0] - draft_model_input = self.make_verify_extend_input( - base_input=main_model_input, - input_ids=next_token_ids, - draft_hidden=target_hidden, - ) - draft_model_output = draft_model.forward(draft_model_input) - - if draft_step == 0: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - schedule_scores=schedule_scores, - ) - - draft_logits = draft_model_output.logits.index_select(0, selected_rows) - if collect_dynamic_probs: - draft_next_token_ids, selected_draft_prob = self._gen_argmax_token_ids_and_prob( - ModelOutput(logits=draft_logits) - ) - schedule_scores[:, 0] = self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=selected_draft_prob, - verify_row_count=verify_row_count, - ) - chain_survival = selected_draft_prob.float().clamp(0.01, 0.99) - else: - draft_next_token_ids = self._gen_argmax_token_ids(ModelOutput(logits=draft_logits)) - chain_survival = None - draft_hidden = draft_model_output.spec_hidden.index_select(0, selected_rows) - proposal_token_ids[selected_rows, 1] = draft_next_token_ids - if draft_step == 1: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - schedule_scores=schedule_scores, - ) - - eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * (draft_step - 1)) - eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) - - selected_seq_len = main_model_input.b_seq_len.index_select(0, selected_rows) + 1 - selected_req_idx = main_model_input.b_req_idx.index_select(0, selected_rows) - selected_mtp_index = torch.zeros_like(selected_req_idx) - selected_position_delta = ( - main_model_input.b_position_delta.index_select(0, selected_rows) - if main_model_input.b_position_delta is not None - else None - ) - one_row_group_marks = torch.ones((num_reqs,), dtype=torch.int32, device=next_token_ids.device) - draft_row_budget = max(1, int(main_model_input.batch_size) - int(num_reqs)) - - for step in range(1, draft_step): - next_depth = step + 1 - active_count = ( - self._get_pruned_active_count( - current_count=int(selected_rows.shape[0]), - draft_row_budget=draft_row_budget, - next_depth=next_depth, - ) - if collect_dynamic_probs - else int(selected_rows.shape[0]) - ) - if active_count < int(selected_rows.shape[0]): - keep_rows = torch.topk( - chain_survival, - k=active_count, - largest=True, - sorted=False, - ).indices - selected_rows = selected_rows.index_select(0, keep_rows) - draft_next_token_ids = draft_next_token_ids.index_select(0, keep_rows) - draft_hidden = draft_hidden.index_select(0, keep_rows) - selected_seq_len = selected_seq_len.index_select(0, keep_rows) - selected_req_idx = selected_req_idx.index_select(0, keep_rows) - selected_mtp_index = selected_mtp_index.index_select(0, keep_rows) - if selected_position_delta is not None: - selected_position_delta = selected_position_delta.index_select(0, keep_rows) - one_row_group_marks = one_row_group_marks.index_select(0, keep_rows) - chain_survival = chain_survival.index_select(0, keep_rows) - - mem_start = (step - 1) * num_reqs - mem_indexes_i = eagle_mem_indexes[mem_start : mem_start + active_count] - draft_input = self.make_single_step_decode_input( - base_input=main_model_input, - input_ids=draft_next_token_ids, - draft_hidden=draft_hidden, - b_req_idx=selected_req_idx, - b_mtp_index=selected_mtp_index, - b_seq_len=selected_seq_len, - b_position_delta=selected_position_delta, - mem_indexes=mem_indexes_i, - b_mark_shared_group=one_row_group_marks, - max_kv_seq_len=main_model_input.max_kv_seq_len + step, - ) - draft_output = draft_model.forward(draft_input) - if collect_dynamic_probs: - draft_next_token_ids, selected_draft_prob = self._gen_argmax_token_ids_and_prob(draft_output) - schedule_scores[:, step] = self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=selected_draft_prob, - verify_row_count=verify_row_count, - ) - chain_survival = chain_survival * selected_draft_prob.float().clamp(0.01, 0.99) - else: - draft_next_token_ids = self._gen_argmax_token_ids(draft_output) - proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids - draft_hidden = draft_output.spec_hidden - selected_seq_len = selected_seq_len + 1 - - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=eagle_mem_indexes_cpu, - schedule_scores=schedule_scores, - ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py index 194b415964..a0ba2557bd 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -8,8 +8,8 @@ from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal -class RecurrentEagleMTPProposer(BaseSpecProposer): - """Shared draft-state setup for recurrent Eagle MTP proposers.""" +class AutoregressiveEagleProposer(BaseSpecProposer): + """Shared autoregressive drafting flow for EAGLE-family proposers.""" def get_draft_steps(self): return tuple(range(1, self.backend.max_draft_step + 1)) @@ -21,11 +21,10 @@ def get_draft_cost_ms( verify_batch_size: int, draft_step: int, ) -> float: - # The mandatory extend processes all verified rows and produces the - # first candidate. Later recurrent forwards process one row per request. - extend_cost_ms = draft_infer_costs.estimate(verify_batch_size) - decode_cost_ms = draft_infer_costs.get(req_num) * (draft_step - 1) - return extend_cost_ms + decode_cost_ms + draft_cost_ms = draft_infer_costs.estimate(verify_batch_size) + if draft_step > 1: + draft_cost_ms += draft_infer_costs.get(req_num) * (draft_step - 1) + return draft_cost_ms def build_draft_state_from_prefill( self, @@ -35,12 +34,12 @@ def build_draft_state_from_prefill( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - draft_model_input = prepare_mtp_prefill_inputs( + prepare_mtp_prefill_inputs( model_input=target_model_input, b_next_token_ids=next_token_ids, mtp_draft_input_hiddens=target_model_output.spec_hidden, ) - self.backend.draft_models[0].forward(draft_model_input) + self.backend.draft_models[0].forward(target_model_input) def build_draft_state_from_prefill_overlap( self, @@ -53,17 +52,17 @@ def build_draft_state_from_prefill_overlap( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - draft_model_input0 = prepare_mtp_prefill_inputs( + prepare_mtp_prefill_inputs( model_input=target_model_input0, b_next_token_ids=next_token_ids0, mtp_draft_input_hiddens=target_model_output0.spec_hidden, ) - draft_model_input1 = prepare_mtp_prefill_inputs( + prepare_mtp_prefill_inputs( model_input=target_model_input1, b_next_token_ids=next_token_ids1, mtp_draft_input_hiddens=target_model_output1.spec_hidden, ) - self.backend.draft_models[0].microbatch_overlap_prefill(draft_model_input0, draft_model_input1) + self.backend.draft_models[0].microbatch_overlap_prefill(target_model_input0, target_model_input1) def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return draft_token_ids @@ -72,76 +71,27 @@ def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: return self._map_draft_token_ids(self.backend._gen_argmax_token_ids(model_output)) def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): - draft_token_ids, draft_probs = self.backend._gen_argmax_token_ids_and_prob(model_output) - return self._map_draft_token_ids(draft_token_ids), draft_probs + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(model_output) + return self._map_draft_token_ids(draft_token_ids), draft_token_probs - def make_verify_extend_input( + def prepare_verify_extend_input( self, - base_input: ModelInput, + model_input: ModelInput, input_ids: torch.Tensor, - draft_hidden: torch.Tensor, - ) -> ModelInput: - new_input = copy.copy(base_input) - batch_size = int(input_ids.shape[0]) - new_input.is_prefill = True - new_input.batch_size = batch_size - new_input.total_token_num = batch_size - new_input.prefix_total_token_num = 0 - new_input.max_q_seq_len = 1 - new_input.max_cache_len = max(0, int(base_input.max_cache_len or 0)) - new_input.input_ids = input_ids - new_input.mtp_draft_input_hiddens = draft_hidden - new_input.mem_indexes_cpu = None - # Treat each verify row as a one-token extend request. This writes the - # accepted prefix into draft KV without giving recurrent decode a - # second attention topology or CUDA graph variant. - new_input.b_ready_cache_len = base_input.b_seq_len - 1 - new_input.b_prefill_start_loc = torch.arange( - batch_size, + target_hidden: torch.Tensor, + ) -> None: + model_input.is_prefill = True + model_input.total_token_num = model_input.batch_size + model_input.prefix_total_token_num = 0 + model_input.max_cache_len = max(0, int(model_input.max_cache_len or 0)) + model_input.input_ids = input_ids + model_input.mtp_draft_input_hiddens = target_hidden + model_input.b_ready_cache_len = model_input.b_seq_len - 1 + model_input.b_prefill_start_loc = torch.arange( + model_input.batch_size, dtype=torch.int32, device=input_ids.device, ) - new_input.b_prefill_has_output_cpu = [False] * batch_size - new_input.b_position_delta = None - new_input.b_is_decode_req = None - new_input.multimodal_params = [{"images": [], "audios": []}] * batch_size - return new_input - - def make_single_step_decode_input( - self, - base_input: ModelInput, - input_ids: torch.Tensor, - draft_hidden: torch.Tensor, - b_req_idx: torch.Tensor, - b_mtp_index: torch.Tensor, - b_seq_len: torch.Tensor, - b_position_delta: torch.Tensor, - mem_indexes: torch.Tensor, - b_mark_shared_group: torch.Tensor, - max_kv_seq_len: int, - ) -> ModelInput: - new_input = copy.copy(base_input) - new_input.batch_size = b_seq_len.shape[0] - new_input.input_ids = input_ids - new_input.mtp_draft_input_hiddens = draft_hidden - new_input.b_req_idx = b_req_idx - new_input.b_mtp_index = b_mtp_index - new_input.b_seq_len = b_seq_len - new_input.b_position_delta = b_position_delta - new_input.mem_indexes = mem_indexes - new_input.mem_indexes_cpu = None - new_input.b_mark_shared_group = b_mark_shared_group - new_input.b_shared_seq_len = None - new_input.max_q_seq_len = 1 - new_input.max_kv_seq_len = max_kv_seq_len - new_input.total_token_num = new_input.batch_size * max_kv_seq_len - new_input.draft_step = 0 - # Recurrent Eagle decode only needs a correctly sized placeholder - # list. Nested per-row allocations otherwise sit between graph - # replays and extend the draft proposal critical path. - empty_multimodal_params = {"images": [], "audios": []} - new_input.multimodal_params = [empty_multimodal_params] * new_input.batch_size - return new_input @staticmethod def _pad_step_mem_indexes( @@ -158,6 +108,104 @@ def _pad_step_mem_indexes( padded[: real_mem_indexes.shape[0]].copy_(real_mem_indexes) return padded + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + """Run the common autoregressive EAGLE proposal flow.""" + + verify_row_count = int(next_token_ids.shape[0]) + request_count = int(b_req_mtp_start_loc.shape[0]) + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + proposal_token_ids = next_token_ids.new_full( + (verify_row_count, draft_step + 1), + fill_value=1, + ) + proposal_token_ids[:, 0].copy_(next_token_ids) + collect_schedule_scores = self.enable_dynamic_spec + schedule_scores = ( + torch.zeros( + (verify_row_count, draft_step), + dtype=torch.float32, + device=next_token_ids.device, + ) + if collect_schedule_scores + else None + ) + + draft_model = self.backend.draft_models[0] + position_delta = main_model_input.b_position_delta + self.prepare_verify_extend_input( + model_input=main_model_input, + input_ids=next_token_ids, + target_hidden=main_model_output.spec_hidden, + ) + extend_output = draft_model.forward(main_model_input) + + accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) + if collect_schedule_scores: + draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(accepted_tail_output) + schedule_scores[accepted_tail_rows, 0] = draft_token_probs.float() + else: + draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) + proposal_token_ids[accepted_tail_rows, 1] = draft_token_ids + draft_hidden = extend_output.spec_hidden.index_select(0, accepted_tail_rows) + + if draft_step == 1: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + schedule_scores=schedule_scores, + ) + + extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(request_count * (draft_step - 1)) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) + draft_seq_lens = main_model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1 + max_kv_seq_len = main_model_input.max_kv_seq_len + draft_input = copy.copy(main_model_input) + draft_input.is_prefill = False + draft_input.batch_size = request_count + draft_input.b_req_idx = main_model_input.b_req_idx.index_select(0, accepted_tail_rows) + draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) + draft_input.b_seq_len = draft_seq_lens + draft_input.b_position_delta = ( + position_delta.index_select(0, accepted_tail_rows) if position_delta is not None else None + ) + draft_input.b_mark_shared_group = torch.ones_like(draft_input.b_req_idx) + draft_input.b_shared_seq_len = None + draft_input.draft_step = 0 + if len(draft_input.multimodal_params) != request_count: + empty_multimodal_params = {"images": [], "audios": []} + draft_input.multimodal_params = [empty_multimodal_params] * request_count + + for step in range(1, draft_step): + mem_start = (step - 1) * request_count + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_input.mem_indexes = extra_mem_indexes[mem_start : mem_start + request_count] + draft_input.max_kv_seq_len = max_kv_seq_len + step + draft_input.total_token_num = request_count * draft_input.max_kv_seq_len + draft_output = draft_model.forward(draft_input) + if collect_schedule_scores: + draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(draft_output) + schedule_scores[accepted_tail_rows, step] = draft_token_probs.float() + else: + draft_token_ids = self._gen_argmax_token_ids(draft_output) + proposal_token_ids[accepted_tail_rows, step + 1] = draft_token_ids + draft_hidden = draft_output.spec_hidden + draft_seq_lens.add_(1) + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, + schedule_scores=schedule_scores, + ) + def propose_next_overlap( self, main_model_input0: ModelInput, @@ -172,30 +220,35 @@ def propose_next_overlap( accept_len1: torch.Tensor, draft_step: int, ) -> SpecProposal: - """Run recurrent Eagle with DP microbatch overlap. + """Run autoregressive EAGLE drafting with DP microbatch overlap. Target verification remains physically padded to ``B * (K + 1)`` for DP collectives. Draft state is corrected once over those verify rows; - recurrent draft decode then runs one row per logical request (plus + autoregressive drafting then runs one row per logical request (plus HOLD rows required to keep the two microbatches shape-compatible). """ verify_width = self.backend.max_draft_step + 1 - inputs = (main_model_input0, main_model_input1) - outputs = (main_model_output0, main_model_output1) - next_ids = (next_token_ids0, next_token_ids1) - real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) - accept_lens = (accept_len0, accept_len1) - - request_capacities = [] - real_request_nums = [] - selected_rows = [] - extend_inputs = [] - for model_input, model_output, token_ids, real_rows, accept_len in zip( - inputs, outputs, next_ids, real_verify_rows, accept_lens + model_inputs = (main_model_input0, main_model_input1) + model_outputs = (main_model_output0, main_model_output1) + next_token_ids_by_batch = (next_token_ids0, next_token_ids1) + real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) + accept_lens_by_batch = (accept_len0, accept_len1) + position_deltas_by_batch = tuple(model_input.b_position_delta for model_input in model_inputs) + max_kv_seq_lens_by_batch = tuple(model_input.max_kv_seq_len for model_input in model_inputs) + + request_capacities_by_batch = [] + real_request_counts = [] + accepted_tail_rows_by_batch = [] + for model_input, model_output, token_ids, real_verify_row_count, accept_len in zip( + model_inputs, + model_outputs, + next_token_ids_by_batch, + real_verify_row_counts, + accept_lens_by_batch, ): request_capacity = model_input.batch_size // verify_width - real_request_num = real_rows // verify_width + real_request_count = real_verify_row_count // verify_width starts = torch.arange( 0, model_input.batch_size, @@ -203,70 +256,45 @@ def propose_next_overlap( dtype=torch.int32, device=token_ids.device, ) - selected = self.select_accepted_tail_rows( - b_req_mtp_start_loc=starts, - accept_len=accept_len, - ) - request_capacities.append(request_capacity) - real_request_nums.append(real_request_num) - selected_rows.append(selected) - extend_inputs.append( - self.make_verify_extend_input( - base_input=model_input, - input_ids=token_ids, - draft_hidden=model_output.spec_hidden, - ) + accepted_tail_rows = (starts + accept_len - 1).long() + request_capacities_by_batch.append(request_capacity) + real_request_counts.append(real_request_count) + accepted_tail_rows_by_batch.append(accepted_tail_rows) + self.prepare_verify_extend_input( + model_input=model_input, + input_ids=token_ids, + target_hidden=model_output.spec_hidden, ) - real_row_count = real_verify_rows0 + real_verify_rows1 + verify_row_count = real_verify_rows0 + real_verify_rows1 proposal_token_ids = next_token_ids0.new_full( - (real_row_count, draft_step + 1), + (verify_row_count, draft_step + 1), fill_value=1, ) proposal_token_ids[:real_verify_rows0, 0].copy_(next_token_ids0[:real_verify_rows0]) proposal_token_ids[real_verify_rows0:, 0].copy_(next_token_ids1[:real_verify_rows1]) draft_model = self.backend.draft_models[0] - extend_outputs = draft_model.microbatch_overlap_prefill(*extend_inputs) - if draft_step == 0: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - ) + extend_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) - draft_next_token_ids = [] - draft_hiddens = [] - selected_seq_lens = [] - selected_req_idxs = [] - selected_mtp_idxs = [] - selected_position_deltas = [] - group_marks = [] + draft_token_ids_by_batch = [] + draft_hiddens_by_batch = [] + draft_seq_lens_by_batch = [] + draft_req_indices_by_batch = [] + proposal_rows_by_batch = [] proposal_row_offsets = (0, real_verify_rows0) - for index, (model_input, extend_output, selected, real_request_num) in enumerate( - zip(inputs, extend_outputs, selected_rows, real_request_nums) + for batch_index, (model_input, extend_output, accepted_tail_rows, real_request_count) in enumerate( + zip(model_inputs, extend_outputs, accepted_tail_rows_by_batch, real_request_counts) ): - selected_output = ModelOutput(logits=extend_output.logits.index_select(0, selected)) - step_token_ids = self._gen_argmax_token_ids(selected_output) - draft_next_token_ids.append(step_token_ids) - draft_hiddens.append(extend_output.spec_hidden.index_select(0, selected)) - selected_seq_lens.append(model_input.b_seq_len.index_select(0, selected) + 1) - selected_req_idxs.append(model_input.b_req_idx.index_select(0, selected)) - selected_mtp_idxs.append(torch.zeros_like(selected_req_idxs[-1])) - selected_position_deltas.append( - model_input.b_position_delta.index_select(0, selected) - if model_input.b_position_delta is not None - else None - ) - group_marks.append( - torch.ones( - request_capacities[index], - dtype=torch.int32, - device=step_token_ids.device, - ) - ) - if real_request_num > 0: - proposal_rows = selected[:real_request_num].to(torch.long) + proposal_row_offsets[index] - proposal_token_ids[proposal_rows, 1] = step_token_ids[:real_request_num] + accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) + draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) + draft_token_ids_by_batch.append(draft_token_ids) + draft_hiddens_by_batch.append(extend_output.spec_hidden.index_select(0, accepted_tail_rows)) + draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) + draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) + proposal_rows = accepted_tail_rows[:real_request_count] + proposal_row_offsets[batch_index] + proposal_rows_by_batch.append(proposal_rows) + proposal_token_ids[proposal_rows, 1] = draft_token_ids[:real_request_count] if draft_step == 1: return SpecProposal( @@ -274,50 +302,56 @@ def propose_next_overlap( extra_mem_indexes_cpu=None, ) - total_real_requests = sum(real_request_nums) - extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(total_real_requests * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) + for batch_index, model_input in enumerate(model_inputs): + model_input.is_prefill = False + model_input.batch_size = request_capacities_by_batch[batch_index] + model_input.b_req_idx = draft_req_indices_by_batch[batch_index] + model_input.b_mtp_index = torch.zeros_like(model_input.b_req_idx) + model_input.b_seq_len = draft_seq_lens_by_batch[batch_index] + model_input.b_position_delta = ( + position_deltas_by_batch[batch_index].index_select(0, accepted_tail_rows_by_batch[batch_index]) + if position_deltas_by_batch[batch_index] is not None + else None + ) + model_input.b_mark_shared_group = torch.ones_like(model_input.b_req_idx) + model_input.b_shared_seq_len = None + model_input.draft_step = 0 + if len(model_input.multimodal_params) != model_input.batch_size: + empty_multimodal_params = {"images": [], "audios": []} + model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size + + total_real_request_count = sum(real_request_counts) + extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(total_real_request_count * (draft_step - 1)) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) hold_mem_index = self.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX for step in range(1, draft_step): - mem_start = (step - 1) * total_real_requests - step_mem_indexes = extra_mem_indexes[mem_start : mem_start + total_real_requests] - step_inputs = [] + mem_start = (step - 1) * total_real_request_count + step_mem_indexes = extra_mem_indexes[mem_start : mem_start + total_real_request_count] real_mem_start = 0 - for index, model_input in enumerate(inputs): - real_request_num = real_request_nums[index] - real_mem_indexes = step_mem_indexes[real_mem_start : real_mem_start + real_request_num] - real_mem_start += real_request_num + for batch_index, model_input in enumerate(model_inputs): + real_request_count = real_request_counts[batch_index] + real_mem_indexes = step_mem_indexes[real_mem_start : real_mem_start + real_request_count] + real_mem_start += real_request_count padded_mem_indexes = self._pad_step_mem_indexes( real_mem_indexes=real_mem_indexes, - request_capacity=request_capacities[index], + request_capacity=request_capacities_by_batch[batch_index], hold_mem_index=hold_mem_index, ) - step_inputs.append( - self.make_single_step_decode_input( - base_input=model_input, - input_ids=draft_next_token_ids[index], - draft_hidden=draft_hiddens[index], - b_req_idx=selected_req_idxs[index], - b_mtp_index=selected_mtp_idxs[index], - b_seq_len=selected_seq_lens[index], - b_position_delta=selected_position_deltas[index], - mem_indexes=padded_mem_indexes, - b_mark_shared_group=group_marks[index], - max_kv_seq_len=model_input.max_kv_seq_len + step, - ) - ) + model_input.input_ids = draft_token_ids_by_batch[batch_index] + model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] + model_input.mem_indexes = padded_mem_indexes + model_input.max_kv_seq_len = max_kv_seq_lens_by_batch[batch_index] + step + model_input.total_token_num = model_input.batch_size * model_input.max_kv_seq_len - step_outputs = draft_model.microbatch_overlap_decode(*step_inputs) - for index, step_output in enumerate(step_outputs): - step_token_ids = self._gen_argmax_token_ids(step_output) - draft_next_token_ids[index] = step_token_ids - draft_hiddens[index] = step_output.spec_hidden - selected_seq_lens[index] = selected_seq_lens[index] + 1 - real_request_num = real_request_nums[index] - if real_request_num > 0: - proposal_rows = selected_rows[index][:real_request_num].to(torch.long) + proposal_row_offsets[index] - proposal_token_ids[proposal_rows, step + 1] = step_token_ids[:real_request_num] + draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + draft_token_ids = self._gen_argmax_token_ids(draft_output) + draft_token_ids_by_batch[batch_index] = draft_token_ids + draft_hiddens_by_batch[batch_index] = draft_output.spec_hidden + draft_seq_lens_by_batch[batch_index].add_(1) + real_request_count = real_request_counts[batch_index] + proposal_token_ids[proposal_rows_by_batch[batch_index], step + 1] = draft_token_ids[:real_request_count] return SpecProposal( token_ids=proposal_token_ids, @@ -325,8 +359,8 @@ def propose_next_overlap( ) -class EagleMTPProposer(RecurrentEagleMTPProposer): - """Recurrent Eagle MTP proposer. +class EagleMTPProposer(AutoregressiveEagleProposer): + """Autoregressive EAGLE proposer backed by an MTP draft model. The draft model keeps a cache and repeatedly feeds the previous proposal hidden back into the same draft model. @@ -351,61 +385,50 @@ def propose_next_overlap( DPEP overlap requires every rank to keep the target verify layout throughout draft decode. Each logical request therefore remains a contiguous group of ``draft_step + 1`` rows; accepted rows are not - compacted into a one-row recurrent batch on this path. + compacted into a one-row autoregressive draft batch on this path. """ verify_width = self.backend.max_draft_step + 1 model_inputs = (main_model_input0, main_model_input1) - real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) - real_request_nums = tuple(row_count // verify_width for row_count in real_verify_rows) - request_capacities = tuple(model_input.batch_size // verify_width for model_input in model_inputs) - total_real_requests = sum(real_request_nums) - - token_id_steps = [ - torch.cat( - [ - next_token_ids0[:real_verify_rows0], - next_token_ids1[:real_verify_rows1], - ], - dim=0, - ) - ] - if draft_step == 0: - return SpecProposal( - token_ids=token_id_steps[0].unsqueeze(1), - extra_mem_indexes_cpu=None, - ) + real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) + real_request_counts = tuple(row_count // verify_width for row_count in real_verify_row_counts) + request_capacities_by_batch = tuple(model_input.batch_size // verify_width for model_input in model_inputs) + total_real_request_count = sum(real_request_counts) - extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(total_real_requests * draft_step) + proposal_token_ids = next_token_ids0.new_empty((sum(real_verify_row_counts), draft_step + 1)) + proposal_token_ids[:real_verify_rows0, 0] = next_token_ids0[:real_verify_rows0] + proposal_token_ids[real_verify_rows0:, 0] = next_token_ids1[:real_verify_rows1] + + extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(total_real_request_count * draft_step) extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) - split = real_request_nums[0] * draft_step - microbatch_mem_indexes = ( + split = real_request_counts[0] * draft_step + extra_mem_indexes_by_batch = ( extra_mem_indexes[:split], extra_mem_indexes[split:], ) - draft_token_ids = [next_token_ids0, next_token_ids1] - draft_hiddens = [main_model_output0.spec_hidden, main_model_output1.spec_hidden] + draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] + draft_hiddens_by_batch = [main_model_output0.spec_hidden, main_model_output1.spec_hidden] draft_model = self.backend.draft_models[0] hold_mem_index = self.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX for step in range(draft_step): - for index, model_input in enumerate(model_inputs): - model_input.input_ids = draft_token_ids[index] - model_input.mtp_draft_input_hiddens = draft_hiddens[index] + for batch_index, model_input in enumerate(model_inputs): + model_input.input_ids = draft_token_ids_by_batch[batch_index] + model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) - for index, (model_input, draft_output) in enumerate(zip(model_inputs, draft_outputs)): + for batch_index, (model_input, draft_output) in enumerate(zip(model_inputs, draft_outputs)): model_input.b_seq_len += 1 model_input.max_kv_seq_len += 1 - real_request_num = real_request_nums[index] - mem_start = step * real_request_num - step_mem_indexes = microbatch_mem_indexes[index][mem_start : mem_start + real_request_num] + real_request_count = real_request_counts[batch_index] + mem_start = step * real_request_count + step_mem_indexes = extra_mem_indexes_by_batch[batch_index][mem_start : mem_start + real_request_count] step_mem_indexes = self._pad_step_mem_indexes( real_mem_indexes=step_mem_indexes, - request_capacity=request_capacities[index], + request_capacity=request_capacities_by_batch[batch_index], hold_mem_index=hold_mem_index, ) model_input.mem_indexes = torch.cat( @@ -416,132 +439,13 @@ def propose_next_overlap( dim=1, ).view(-1) - draft_token_ids[index] = self._gen_argmax_token_ids(draft_output) - draft_hiddens[index] = draft_output.spec_hidden + draft_token_ids_by_batch[batch_index] = self._gen_argmax_token_ids(draft_output) + draft_hiddens_by_batch[batch_index] = draft_output.spec_hidden - token_id_steps.append( - torch.cat( - [ - draft_token_ids[0][:real_verify_rows0], - draft_token_ids[1][:real_verify_rows1], - ], - dim=0, - ) - ) - - return SpecProposal( - token_ids=torch.stack(token_id_steps, dim=1), - extra_mem_indexes_cpu=extra_mem_indexes_cpu, - ) - - def propose_next( - self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> SpecProposal: - verify_row_count = int(next_token_ids.shape[0]) - num_reqs = int(b_req_mtp_start_loc.shape[0]) - selected_rows = self.select_accepted_tail_rows( - b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=accept_len, - ) - proposal_token_ids = next_token_ids.new_full( - (verify_row_count, draft_step + 1), - fill_value=1, - ) - proposal_token_ids[:, 0].copy_(next_token_ids) - schedule_scores = ( - torch.zeros( - (verify_row_count, draft_step), - dtype=torch.float32, - device=next_token_ids.device, - ) - if self.enable_dynamic_spec - else None - ) - - draft_model = self.backend.draft_models[0] - extend_input = self.make_verify_extend_input( - base_input=main_model_input, - input_ids=next_token_ids, - draft_hidden=main_model_output.spec_hidden, - ) - extend_output = draft_model.forward(extend_input) - - if draft_step == 0: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - schedule_scores=schedule_scores, - ) - - selected_logits = extend_output.logits.index_select(0, selected_rows) - selected_output = ModelOutput(logits=selected_logits) - if self.enable_dynamic_spec: - draft_next_token_ids, selected_prob = self._gen_argmax_token_ids_and_prob(selected_output) - schedule_scores[:, 0] = self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=selected_prob, - verify_row_count=verify_row_count, - ) - else: - draft_next_token_ids = self._gen_argmax_token_ids(selected_output) - proposal_token_ids[selected_rows, 1] = draft_next_token_ids - draft_hidden = extend_output.spec_hidden.index_select(0, selected_rows) - - if draft_step == 1: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - schedule_scores=schedule_scores, - ) - - eagle_mem_indexes_cpu = self.alloc_extra_mem_indexes(num_reqs * (draft_step - 1)) - eagle_mem_indexes = eagle_mem_indexes_cpu.cuda(non_blocking=True) - selected_seq_len = main_model_input.b_seq_len.index_select(0, selected_rows) + 1 - selected_req_idx = main_model_input.b_req_idx.index_select(0, selected_rows) - selected_mtp_index = torch.zeros_like(selected_req_idx) - selected_position_delta = ( - main_model_input.b_position_delta.index_select(0, selected_rows) - if main_model_input.b_position_delta is not None - else None - ) - one_row_group_marks = torch.ones(num_reqs, dtype=torch.int32, device=next_token_ids.device) - - for step in range(1, draft_step): - mem_start = (step - 1) * num_reqs - draft_input = self.make_single_step_decode_input( - base_input=main_model_input, - input_ids=draft_next_token_ids, - draft_hidden=draft_hidden, - b_req_idx=selected_req_idx, - b_mtp_index=selected_mtp_index, - b_seq_len=selected_seq_len, - b_position_delta=selected_position_delta, - mem_indexes=eagle_mem_indexes[mem_start : mem_start + num_reqs], - b_mark_shared_group=one_row_group_marks, - max_kv_seq_len=main_model_input.max_kv_seq_len + step, - ) - draft_output = draft_model.forward(draft_input) - if self.enable_dynamic_spec: - draft_next_token_ids, selected_prob = self._gen_argmax_token_ids_and_prob(draft_output) - schedule_scores[:, step] = self.scatter_selected_step_probs( - selected_rows=selected_rows, - selected_probs=selected_prob, - verify_row_count=verify_row_count, - ) - else: - draft_next_token_ids = self._gen_argmax_token_ids(draft_output) - proposal_token_ids[selected_rows, step + 1] = draft_next_token_ids - draft_hidden = draft_output.spec_hidden - selected_seq_len = selected_seq_len + 1 + proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids_by_batch[0][:real_verify_rows0] + proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids_by_batch[1][:real_verify_rows1] return SpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=eagle_mem_indexes_cpu, - schedule_scores=schedule_scores, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py b/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py new file mode 100644 index 0000000000..3f75e1b326 --- /dev/null +++ b/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +import copy + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer + + +class ParallelBlockProposer(BaseSpecProposer): + """Shared state and input preparation for parallel block drafters. + + DFlash and DSpark consume verified target output, extend the drafter KV + cache with target hidden rows, then generate a new block in one parallel + backbone forward. + """ + + def get_draft_steps(self): + return (self.backend.max_draft_step,) + + def get_draft_cost_ms( + self, + draft_infer_costs, + req_num: int, + verify_batch_size: int, + draft_step: int, + ) -> float: + block_size = self.backend.draft_models[0].block_size + # One forward commits verified rows; another generates one complete + # checkpoint-defined block per request. + extend_cost_ms = draft_infer_costs.estimate(verify_batch_size) + block_cost_ms = draft_infer_costs.get(req_num * block_size) + return extend_cost_ms + block_cost_ms + + @torch.no_grad() + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + target_hidden = target_model_output.spec_hidden + if target_hidden.numel() == 0: + return + + # Parallel block drafters consume target hidden states directly. + target_model_input.mtp_draft_input_hiddens = target_hidden + self.backend.draft_models[0].forward(target_model_input) + + def extend_draft_kv_cache(self, main_model_input: ModelInput, target_hidden: torch.Tensor) -> None: + # Target decode has finished; reuse its row-aligned input for draft KV commit. + main_model_input.total_token_num = main_model_input.batch_size + main_model_input.prefix_total_token_num = 0 + main_model_input.is_prefill = True + main_model_input.b_ready_cache_len = main_model_input.b_seq_len - 1 + main_model_input.b_prefill_start_loc = torch.arange( + main_model_input.batch_size, + dtype=torch.int32, + device=target_hidden.device, + ) + main_model_input.mtp_draft_input_hiddens = target_hidden + self.backend.draft_models[0].forward(main_model_input) + + def build_block_draft_input( + self, + main_model_input: ModelInput, + next_token_ids: torch.Tensor, + accepted_tail_rows: torch.Tensor, + request_count: int, + ): + draft_model = self.backend.draft_models[0] + block_size = int(draft_model.block_size) + extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(request_count * block_size) + + block_input_ids = next_token_ids.new_full( + (request_count * block_size,), + fill_value=draft_model.mask_token_id, + ) + # Block input layout: [accepted token, mask, ..., mask]. The block + # logits become [base token + draft tokens] for target verification. + block_input_ids[::block_size] = next_token_ids.index_select(0, accepted_tail_rows) + + block_offsets = torch.arange( + block_size, + dtype=main_model_input.b_seq_len.dtype, + device=next_token_ids.device, + ) + draft_input = copy.copy(main_model_input) + draft_input.input_ids = block_input_ids + draft_input.total_token_num = draft_input.input_ids.shape[0] + draft_input.batch_size = draft_input.total_token_num + draft_input.max_q_seq_len = 1 + draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size + draft_input.draft_step = block_size - 1 + draft_input.b_req_idx = ( + main_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() + ) + draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) + # copy_kv_index_to_req and FA3 use these lengths to place scratch KV and + # compute the block cache length. + draft_input.b_seq_len = ( + (main_model_input.b_seq_len.index_select(0, accepted_tail_rows)[:, None] + block_offsets[None, :] + 1) + .reshape(-1) + .contiguous() + ) + # Position delta is request-level metadata shared by every block row. + draft_input.b_position_delta = ( + main_model_input.b_position_delta.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() + ) + draft_input.mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) + draft_input.b_mark_shared_group = torch.zeros_like(draft_input.b_req_idx) + draft_input.b_mark_shared_group[block_size - 1 :: block_size] = block_size + empty_multimodal_params = {"images": [], "audios": []} + draft_input.multimodal_params = [empty_multimodal_params] * draft_input.batch_size + return draft_input, extra_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py index 5ad90cd29c..c044e53874 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py @@ -9,14 +9,8 @@ class VanillaMTPProposer(BaseSpecProposer): """Chained MTP proposer. - This path uses `max_draft_step` independent draft modules. Step i consumes the - hidden feature produced by step i - 1 and predicts one candidate token. - - Target -> draft transfer: - - target prefill/decode captures final hidden states with shape - [token_num, hidden_size] - - SpecEngine injects them into ModelInput.mtp_draft_input_hiddens before each - draft forward + Each draft depth uses an independent MTP module. Module i consumes the + hidden state produced by module i - 1 and predicts the next candidate. """ def get_draft_steps(self): @@ -29,7 +23,7 @@ def get_draft_cost_ms( verify_batch_size: int, draft_step: int, ) -> float: - # Each selected MTP module processes the complete target verify batch. + # Every selected MTP module processes the complete verify batch. return draft_infer_costs.get(verify_batch_size) * draft_step def build_draft_state_from_prefill( @@ -40,17 +34,17 @@ def build_draft_state_from_prefill( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - draft_model_input = target_model_input - source_model_output = target_model_output - draft_next_token_ids = next_token_ids + draft_hidden = target_model_output.spec_hidden + draft_token_ids = next_token_ids for draft_model in self.backend.draft_models: - draft_model_input = prepare_mtp_prefill_inputs( - model_input=draft_model_input, - b_next_token_ids=draft_next_token_ids, - mtp_draft_input_hiddens=source_model_output.spec_hidden, + prepare_mtp_prefill_inputs( + model_input=target_model_input, + b_next_token_ids=draft_token_ids, + mtp_draft_input_hiddens=draft_hidden, ) - source_model_output = draft_model.forward(draft_model_input) - draft_next_token_ids = self.backend._gen_argmax_token_ids(source_model_output) + draft_output = draft_model.forward(target_model_input) + draft_hidden = draft_output.spec_hidden + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) def build_draft_state_from_prefill_overlap( self, @@ -63,30 +57,27 @@ def build_draft_state_from_prefill_overlap( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - draft_model_input0 = target_model_input0 - draft_model_input1 = target_model_input1 - source_model_output0 = target_model_output0 - source_model_output1 = target_model_output1 - draft_next_token_ids0 = next_token_ids0 - draft_next_token_ids1 = next_token_ids1 + draft_hiddens_by_batch = [target_model_output0.spec_hidden, target_model_output1.spec_hidden] + draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] for draft_model in self.backend.draft_models: - draft_model_input0 = prepare_mtp_prefill_inputs( - model_input=draft_model_input0, - b_next_token_ids=draft_next_token_ids0, - mtp_draft_input_hiddens=source_model_output0.spec_hidden, + prepare_mtp_prefill_inputs( + model_input=target_model_input0, + b_next_token_ids=draft_token_ids_by_batch[0], + mtp_draft_input_hiddens=draft_hiddens_by_batch[0], ) - draft_model_input1 = prepare_mtp_prefill_inputs( - model_input=draft_model_input1, - b_next_token_ids=draft_next_token_ids1, - mtp_draft_input_hiddens=source_model_output1.spec_hidden, + prepare_mtp_prefill_inputs( + model_input=target_model_input1, + b_next_token_ids=draft_token_ids_by_batch[1], + mtp_draft_input_hiddens=draft_hiddens_by_batch[1], ) - source_model_output0, source_model_output1 = draft_model.microbatch_overlap_prefill( - draft_model_input0, - draft_model_input1, + draft_outputs = draft_model.microbatch_overlap_prefill( + target_model_input0, + target_model_input1, ) - draft_next_token_ids0 = self.backend._gen_argmax_token_ids(source_model_output0) - draft_next_token_ids1 = self.backend._gen_argmax_token_ids(source_model_output1) + for batch_index, draft_output in enumerate(draft_outputs): + draft_hiddens_by_batch[batch_index] = draft_output.spec_hidden + draft_token_ids_by_batch[batch_index] = self.backend._gen_argmax_token_ids(draft_output) def propose_next( self, @@ -97,13 +88,14 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> SpecProposal: - draft_model_input = main_model_input - draft_next_token_ids = next_token_ids - draft_hidden = main_model_output.spec_hidden if draft_step > 0 else None - all_next_token_ids = [next_token_ids] + verify_row_count = int(next_token_ids.shape[0]) + draft_token_ids = next_token_ids + draft_hidden = main_model_output.spec_hidden + proposal_token_ids = next_token_ids.new_empty((verify_row_count, draft_step + 1)) + proposal_token_ids[:, 0] = next_token_ids schedule_scores = ( torch.empty( - (next_token_ids.shape[0], draft_step), + (verify_row_count, draft_step), dtype=torch.float32, device=next_token_ids.device, ) @@ -113,19 +105,19 @@ def propose_next( for step in range(draft_step): draft_model = self.backend.draft_models[step] - draft_model_input.input_ids = draft_next_token_ids - draft_model_input.mtp_draft_input_hiddens = draft_hidden - draft_model_output = draft_model.forward(draft_model_input) - draft_hidden = draft_model_output.spec_hidden + main_model_input.input_ids = draft_token_ids + main_model_input.mtp_draft_input_hiddens = draft_hidden + draft_output = draft_model.forward(main_model_input) + draft_hidden = draft_output.spec_hidden if self.enable_dynamic_spec: - draft_next_token_ids, draft_prob = self.backend._gen_argmax_token_ids_and_prob(draft_model_output) - schedule_scores[:, step] = draft_prob + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + schedule_scores[:, step] = draft_token_probs else: - draft_next_token_ids = self.backend._gen_argmax_token_ids(draft_model_output) - all_next_token_ids.append(draft_next_token_ids) + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + proposal_token_ids[:, step + 1] = draft_token_ids return SpecProposal( - token_ids=torch.stack(all_next_token_ids, dim=1), + token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, schedule_scores=schedule_scores, ) diff --git a/unit_tests/common/basemodel/test_hidden_collector.py b/unit_tests/common/basemodel/test_hidden_collector.py index dbd311ee60..a352d15bcf 100644 --- a/unit_tests/common/basemodel/test_hidden_collector.py +++ b/unit_tests/common/basemodel/test_hidden_collector.py @@ -37,11 +37,11 @@ def test_hidden_collector_selects_implementation(): def test_draft_hidden_collector_follows_spec_mode(): model = SimpleNamespace(is_mtp_draft_model=True) - recurrent_collector = HiddenCollector(model=model, spec_mode="eagle3") - block_collector = HiddenCollector(model=model, spec_mode="dspark") + autoregressive_collector = HiddenCollector(model=model, spec_mode="eagle3") + parallel_block_collector = HiddenCollector(model=model, spec_mode="dspark") - assert isinstance(recurrent_collector.collectors[0], FinalHiddenCollector) - assert isinstance(block_collector.collectors[0], NoopHiddenCollector) + assert isinstance(autoregressive_collector.collectors[0], FinalHiddenCollector) + assert isinstance(parallel_block_collector.collectors[0], NoopHiddenCollector) def test_hidden_collector_supports_single_and_overlap_forward(): diff --git a/unit_tests/models/qwen2_vl/test_infer_struct.py b/unit_tests/models/qwen2_vl/test_infer_struct.py new file mode 100644 index 0000000000..e62b188940 --- /dev/null +++ b/unit_tests/models/qwen2_vl/test_infer_struct.py @@ -0,0 +1,29 @@ +from types import SimpleNamespace + +import torch + +from lightllm.common.basemodel.infer_struct import InferStateInfo +from lightllm.models.qwen2_vl.infer_struct import Qwen2VLInferStateInfo + + +def test_draft_commit_prefill_uses_existing_position_delta(monkeypatch): + def init_single_token_prefill(self, model): + self.position_ids = torch.tensor([9, 19], dtype=torch.int32) + self.b_q_seq_len = torch.ones(2, dtype=torch.int32) + + monkeypatch.setattr(InferStateInfo, "init_some_extra_state", init_single_token_prefill) + + infer_state = Qwen2VLInferStateInfo() + infer_state.is_prefill = True + infer_state.b_position_delta = torch.tensor([3, 5], dtype=torch.int32) + infer_state.multimodal_params = [{"images": [], "audios": []}] * 2 + model = SimpleNamespace( + config={"rope_scaling": {}}, + _cos_cached=torch.arange(32, dtype=torch.float32).view(32, 1), + _sin_cached=torch.arange(32, dtype=torch.float32).view(32, 1), + ) + + infer_state.init_some_extra_state(model) + + expected_position_ids = torch.tensor([[12, 24]] * 3, dtype=torch.int32) + assert torch.equal(infer_state.position_ids, expected_position_ids) diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index 2b67c64d05..f621e6f7bf 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -119,7 +119,7 @@ def test_fixed_dspark_does_not_require_confidence_head(monkeypatch): model._verify_params() -def test_qwen35_dspark_uses_training_rope_layout(monkeypatch): +def test_qwen35_dspark_adapter_uses_current_checkpoint_rope_layout(monkeypatch): def init_dspark_config(self): self.config = { "dflash_config": {"mask_token_id": 1}, diff --git a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py index b30c01dfb3..d927fb2678 100644 --- a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py @@ -4,16 +4,22 @@ from lightllm.common.basemodel.batch_objs import ModelOutput from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer +from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import ( + AutoregressiveEagleProposer, + EagleMTPProposer, +) class _DraftModel: def __init__(self): self.extend_batch_sizes = None + self.extend_inputs = None self.decode_batch_sizes = [] + self.decode_inputs = [] def microbatch_overlap_prefill(self, input0, input1): self.extend_batch_sizes = (input0.batch_size, input1.batch_size) + self.extend_inputs = (input0, input1) return tuple( ModelOutput( logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), @@ -24,6 +30,7 @@ def microbatch_overlap_prefill(self, input0, input1): def microbatch_overlap_decode(self, input0, input1): self.decode_batch_sizes.append((input0.batch_size, input1.batch_size)) + self.decode_inputs.append((input0, input1)) return tuple( ModelOutput( logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), @@ -36,12 +43,19 @@ def microbatch_overlap_decode(self, input0, input1): def _target_input(batch_size): return SimpleNamespace( batch_size=batch_size, + total_token_num=batch_size, + prefix_total_token_num=None, + input_ids=torch.arange(batch_size, dtype=torch.int64), b_seq_len=torch.arange(batch_size, dtype=torch.int32) + 4, b_req_idx=torch.arange(batch_size, dtype=torch.int32), + b_mtp_index=torch.zeros(batch_size, dtype=torch.int32), b_position_delta=torch.zeros(batch_size, dtype=torch.int32), mem_indexes=torch.arange(batch_size, dtype=torch.int32), + mem_indexes_cpu=torch.arange(batch_size, dtype=torch.int32), max_kv_seq_len=16, max_cache_len=16, + is_prefill=False, + multimodal_params=[{"images": [], "audios": []}] * batch_size, ) @@ -88,6 +102,48 @@ def test_overlap_eagle_keeps_fixed_verify_layout(): assert torch.equal(model_input1.mem_indexes, torch.tensor([2, 2, 4, 5, 3, 5], dtype=torch.int32)) +def test_autoregressive_eagle_reuses_overlap_inputs(): + draft_model = _DraftModel() + backend = SimpleNamespace( + max_draft_step=2, + draft_models=[draft_model], + model=SimpleNamespace( + req_manager=SimpleNamespace( + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + ) + ), + _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), + ) + proposer = AutoregressiveEagleProposer(backend=backend, enable_dynamic_spec=False) + proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) + model_input0 = _target_input(batch_size=6) + model_input1 = _target_input(batch_size=6) + + proposal = proposer.propose_next_overlap( + main_model_input0=model_input0, + main_model_output0=ModelOutput(logits=torch.empty((6, 1)), spec_hidden=torch.ones((6, 2))), + next_token_ids0=torch.arange(6, dtype=torch.int64), + real_verify_rows0=3, + accept_len0=torch.tensor([2, 1], dtype=torch.int32), + main_model_input1=model_input1, + main_model_output1=ModelOutput(logits=torch.empty((6, 1)), spec_hidden=torch.ones((6, 2))), + next_token_ids1=torch.arange(10, 16, dtype=torch.int64), + real_verify_rows1=6, + accept_len1=torch.tensor([1, 3], dtype=torch.int32), + draft_step=2, + ) + + assert draft_model.extend_inputs[0] is model_input0 + assert draft_model.extend_inputs[1] is model_input1 + assert len(draft_model.decode_inputs) == 1 + assert draft_model.decode_inputs[0][0] is model_input0 + assert draft_model.decode_inputs[0][1] is model_input1 + assert draft_model.extend_batch_sizes == (6, 6) + assert draft_model.decode_batch_sizes == [(2, 2)] + assert proposal.token_ids.shape == (9, 3) + assert torch.equal(proposal.extra_mem_indexes_cpu, torch.arange(3, dtype=torch.int32)) + + def test_eagle3_maps_draft_token_ids_in_proposer(): proposer = Eagle3Proposer.__new__(Eagle3Proposer) proposer.backend = SimpleNamespace( diff --git a/unit_tests/server/router/model_infer/speculative/test_planner.py b/unit_tests/server/router/model_infer/speculative/test_planner.py index 70d110ee77..920ada9a14 100644 --- a/unit_tests/server/router/model_infer/speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/speculative/test_planner.py @@ -9,12 +9,17 @@ FixedSpecPlanner, LightSpecPlanner, SpecDecodePlan, + _InferCostMsTable, ) from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer +from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import ( + AutoregressiveEagleProposer, + EagleMTPProposer, +) +from lightllm.server.router.model_infer.speculative.proposers.parallel_block import ParallelBlockProposer from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer @@ -78,6 +83,14 @@ def test_fixed_planner_returns_static_plan(): assert not plan.skip_verify_sync +def test_infer_cost_candidates_include_feasible_boundaries(): + costs = _InferCostMsTable() + costs.update(batch_size=4, infer_cost_ms=1.0) + costs.update(batch_size=8, infer_cost_ms=2.0) + + assert costs.get_batch_size_keys_between(5, 10) == [5, 8, 10] + + def test_engine_routes_only_dspark_to_the_confidence_planner(): assert isinstance(build_planner("eagle3", enable_dynamic_spec=False), FixedSpecPlanner) assert isinstance(build_planner("dspark"), DSparkPlanner) @@ -91,6 +104,14 @@ def test_engine_routes_only_dspark_to_the_confidence_planner(): assert eagle_planner.draft_steps == (1, 2, 3) +def test_proposer_families_use_drafter_standard_abstractions(): + assert issubclass(EagleMTPProposer, AutoregressiveEagleProposer) + assert issubclass(Eagle3Proposer, AutoregressiveEagleProposer) + assert issubclass(DFlashProposer, ParallelBlockProposer) + assert issubclass(DSparkProposer, ParallelBlockProposer) + assert not issubclass(DSparkProposer, DFlashProposer) + + def test_dynamic_plan_filters_selected_rows(): plan = SpecDecodePlan(dynamic_batch_size=2, draft_step=3, pre_draft_step=3) reqs = ["req0", "req0", "req1", "req1"] @@ -312,22 +333,23 @@ def test_eagle_cost_provider_prices_extend_and_decode_separately(): assert draft_cost_ms == 1.5 -def test_eagle3_cost_provider_accounts_for_pruned_recurrent_rows(): - planner = build_lightspec_planner( - max_draft_step=7, - proposer_class=Eagle3Proposer, - ) - for batch_size in (2, 4, 8, 16): - planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) - - draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=planner.draft_infer_costs, - req_num=8, - verify_batch_size=16, - draft_step=7, - ) +def test_autoregressive_eagle_cost_provider_prices_extend_and_decode_rows(): + for proposer_class in (EagleMTPProposer, Eagle3Proposer): + planner = build_lightspec_planner( + max_draft_step=7, + proposer_class=proposer_class, + ) + for batch_size in (2, 4, 8, 16): + planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) + + draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=planner.draft_infer_costs, + req_num=8, + verify_batch_size=16, + draft_step=7, + ) - assert draft_cost_ms == 42.0 + assert draft_cost_ms == 64.0 def test_block_cost_provider_prices_commit_and_complete_block(): @@ -365,7 +387,7 @@ def test_lightspec_records_one_batch_observation_per_configuration(): ) assert planner.progress_ema_by_config[(2, 4, 1)].get() == 1.0 - assert np.isclose(planner.progress_ema_by_config[(2, 4, 3)].get(), 0.95) + assert planner.progress_ema_by_config[(2, 4, 3)].get() == 0.5 assert planner.progress_ema_by_config[(2, 4, 1)].get_count() == 1 assert planner.progress_ema_by_config[(2, 4, 3)].get_count() == 1 @@ -383,7 +405,7 @@ def test_lightspec_high_concurrency_does_not_multiply_ema_updates(): assert planner.progress_ema_by_config[(128, 128, 1)].get_count() == 1 -def test_lightspec_transfers_committed_progress_to_a_nearby_verify_shape(): +def test_lightspec_estimates_unseen_shapes_from_prefix_survival(): planner = build_lightspec_planner() planner.update_verified_batch( accept_lengths=[2, 1], @@ -397,6 +419,38 @@ def test_lightspec_transfers_committed_progress_to_a_nearby_verify_shape(): assert planner._estimate_progress(req_num=2, dynamic_batch_size=4, draft_step=2) == 0.75 +def test_lightspec_does_not_transfer_deep_progress_to_short_drafts(): + planner = build_lightspec_planner() + planner.update_verified_batch( + accept_lengths=[4, 2], + req_num=2, + dynamic_batch_size=8, + verified_draft_step=3, + ) + + assert np.isclose(planner._estimate_progress(req_num=2, dynamic_batch_size=6, draft_step=2), 5 / 6) + + +def test_lightspec_short_current_proposal_can_recover_to_a_deeper_draft(): + planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + for batch_size, target_cost in ((2, 1.0), (4, 1.1), (6, 1.2), (8, 1.3)): + planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) + planner.update_infer_cost(batch_size, 0.01, is_draft_model=True) + planner.update_verified_batch( + accept_lengths=[4, 4], + req_num=2, + dynamic_batch_size=8, + verified_draft_step=3, + ) + planner.pre_draft_step = 1 + + plan = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + + assert plan.dynamic_batch_size <= 4 + assert plan.pre_draft_step == 1 + assert plan.draft_step == 3 + + def test_engine_skips_feedback_for_a_mixed_proposal_batch(): engine = SpecEngine.__new__(SpecEngine) engine.enable_dynamic_spec = True diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index c85c355bcd..979396f6a5 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -150,9 +150,9 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): block_size = 4 max_draft_step = 3 verify_row_count = 5 - selected_rows = torch.tensor([0, 3]) - flat_token_ids = torch.arange(2 * block_size) - flat_token_probs = torch.arange(2 * block_size, dtype=torch.float32) / 10 + accepted_tail_rows = torch.tensor([0, 3]) + flat_draft_token_ids = torch.arange(2 * block_size) + flat_draft_token_probs = torch.arange(2 * block_size, dtype=torch.float32) / 10 draft_model = SimpleNamespace( block_size=block_size, @@ -161,11 +161,10 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): backend = SimpleNamespace( max_draft_step=max_draft_step, draft_models=[draft_model], - _gen_argmax_token_ids_and_prob=lambda _: (flat_token_ids, flat_token_probs), + _gen_argmax_token_ids_and_prob=lambda _: (flat_draft_token_ids, flat_draft_token_probs), ) proposer = DFlashProposer(backend=backend, enable_dynamic_spec=True) proposer.extend_draft_kv_cache = lambda **_: None - proposer.select_accepted_tail_rows = lambda **_: selected_rows proposer.build_block_draft_input = lambda **_: (SimpleNamespace(), torch.tensor([10, 11])) proposal = proposer.propose_next( @@ -177,16 +176,83 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): accept_len=torch.tensor([1, 1]), ) - expected_blocks = flat_token_ids.reshape(2, block_size)[:, :2] - torch.testing.assert_close(proposal.token_ids[selected_rows, 1:], expected_blocks) + expected_blocks = flat_draft_token_ids.reshape(2, block_size)[:, :2] + torch.testing.assert_close(proposal.token_ids[accepted_tail_rows, 1:], expected_blocks) assert proposal.schedule_scores.shape == (verify_row_count, 2) for step in range(2): scores = proposal.schedule_scores[:, step] expected_scores = torch.zeros(verify_row_count) - expected_scores[selected_rows] = flat_token_probs.reshape(2, block_size)[:, step] + expected_scores[accepted_tail_rows] = flat_draft_token_probs.reshape(2, block_size)[:, step] torch.testing.assert_close(scores, expected_scores) +def test_dflash_reuses_decode_input_for_kv_commit(): + forwarded_inputs = [] + draft_model = SimpleNamespace(forward=forwarded_inputs.append) + proposer = DFlashProposer( + backend=SimpleNamespace(draft_models=[draft_model]), + enable_dynamic_spec=False, + ) + input_ids = torch.arange(4) + multimodal_params = [{"images": [], "audios": []}] * 4 + position_delta = torch.zeros(4, dtype=torch.int32) + model_input = SimpleNamespace( + batch_size=4, + total_token_num=20, + prefix_total_token_num=None, + input_ids=input_ids, + b_seq_len=torch.tensor([10, 11, 12, 13], dtype=torch.int32), + b_position_delta=position_delta, + multimodal_params=multimodal_params, + is_prefill=False, + ) + target_hidden = torch.empty(4, 16) + + proposer.extend_draft_kv_cache(model_input, target_hidden) + + assert len(forwarded_inputs) == 1 + assert forwarded_inputs[0] is model_input + assert model_input.input_ids is input_ids + assert model_input.multimodal_params is multimodal_params + assert model_input.total_token_num == model_input.batch_size + assert model_input.prefix_total_token_num == 0 + assert model_input.is_prefill + torch.testing.assert_close(model_input.b_ready_cache_len, model_input.b_seq_len - 1) + torch.testing.assert_close(model_input.b_prefill_start_loc, torch.arange(4, dtype=torch.int32)) + assert model_input.b_position_delta is position_delta + assert model_input.mtp_draft_input_hiddens is target_hidden + + +def test_dflash_expands_position_delta_with_request_block_rows(): + block_size = 3 + proposer = DFlashProposer( + backend=SimpleNamespace( + draft_models=[SimpleNamespace(block_size=block_size, mask_token_id=99)], + ), + enable_dynamic_spec=False, + ) + proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) + model_input = SimpleNamespace( + b_req_idx=torch.tensor([10, 10, 11, 12, 12], dtype=torch.int32), + b_seq_len=torch.tensor([4, 5, 7, 8, 9], dtype=torch.int32), + b_position_delta=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32), + max_kv_seq_len=9, + ) + + draft_input, _ = proposer.build_block_draft_input( + main_model_input=model_input, + next_token_ids=torch.arange(5, dtype=torch.int64), + accepted_tail_rows=torch.tensor([1, 4]), + request_count=2, + ) + + assert torch.equal( + draft_input.b_position_delta, + torch.tensor([11, 11, 11, 21, 21, 21], dtype=torch.int32), + ) + assert torch.equal(draft_input.b_seq_len, torch.tensor([6, 7, 8, 10, 11, 12], dtype=torch.int32)) + + def test_hidden_collector_reads_target_layer_ids(monkeypatch): monkeypatch.setattr( hidden_collector_module.PretrainedConfig, From ee0ab343026d5d2288760df6a9d4b81136dbfe4e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 14 Aug 2026 07:13:13 +0000 Subject: [PATCH 011/103] refactor: centralize MTP decode batch layout --- lightllm/common/basemodel/basemodel.py | 4 +- lightllm/common/basemodel/mtp_manager.py | 49 +++++++++++++++++++ .../model_infer/mode_backend/base_backend.py | 12 ----- .../common/basemodel/test_mtp_manager.py | 47 ++++++++++++++++++ 4 files changed, 99 insertions(+), 13 deletions(-) create mode 100644 lightllm/common/basemodel/mtp_manager.py create mode 100644 unit_tests/common/basemodel/test_mtp_manager.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 5ecd322351..a2c9553c35 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -32,6 +32,7 @@ HiddenCollector, unpad_collected_hidden, ) +from lightllm.common.basemodel.mtp_manager import MtpManager from lightllm.utils.custom_kernel_utis import pad2dim_tensor_to_new_batch from lightllm.utils.envs_utils import ( set_model_init_status, @@ -85,13 +86,14 @@ def __init__(self, kvargs): self.return_all_prompt_logics = kvargs.get("return_all_prompt_logics", False) assert not (self.is_token_healing and self.return_all_prompt_logics), "can not be true in same time" self.data_type = get_llm_data_type() - self.decode_batch_multiplier = kvargs.get("decode_batch_multiplier", 1) self.graph_max_batch_size = kvargs.get("graph_max_batch_size", 16) self.graph_max_batch_size = ( self.graph_max_batch_size // 2 if get_env_start_args().enable_decode_microbatch_overlap else self.graph_max_batch_size ) + self.mtp_manager = MtpManager.get_instance() + self.decode_batch_multiplier = self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) self.graph_max_batch_size = self.graph_max_batch_size * self.decode_batch_multiplier self.graph_max_len_in_batch = kvargs.get("graph_max_len_in_batch", 8192) diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py new file mode 100644 index 0000000000..c914eda891 --- /dev/null +++ b/lightllm/common/basemodel/mtp_manager.py @@ -0,0 +1,49 @@ +from typing import ClassVar, Optional + +from lightllm.utils.envs_utils import get_env_start_args + + +class MtpManager: + """Manage the physical decode layout for main and draft models.""" + + _instance: ClassVar[Optional["MtpManager"]] = None + _CHAINED_DRAFT_MODES = ("vanilla_with_att", "vanilla_no_att") + _RECURRENT_DRAFT_MODES = ("eagle_with_att", "eagle_no_att", "eagle3") + _BLOCK_DRAFT_MODES = ("dspark", "dflash") + + @classmethod + def get_instance(cls) -> "MtpManager": + if cls._instance is None: + cls._instance = cls() + return cls._instance + + def __init__(self): + self.args = get_env_start_args() + + def get_decode_batch_multiplier(self, is_draft_model: bool) -> int: + """Return the physical decode rows used by one logical request.""" + + spec_mode = self.args.mtp_mode + if spec_mode is None: + return 1 + + verify_width = self.args.mtp_step + 1 + + # The main model verifies one target token plus mtp_step draft tokens + # for every logical request, regardless of how the draft is produced. + if not is_draft_model: + return verify_width + + # Chained MTP runs every draft module over the expanded verify layout. + if spec_mode in self._CHAINED_DRAFT_MODES: + return verify_width + + # Recurrent EAGLE draft models decode one row per logical request. + if spec_mode in self._RECURRENT_DRAFT_MODES: + return 1 + + # Block draft models decode mtp_step rows per logical request. + if spec_mode in self._BLOCK_DRAFT_MODES: + return self.args.mtp_step + + return 1 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 91c60f3f9f..e8308af36b 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -126,7 +126,6 @@ def init_model(self, kvargs): model_cfg, _ = PretrainedConfig.get_config_dict(self.weight_dir) - target_decode_batch_multiplier = self.args.mtp_step + 1 if self.args.mtp_mode is not None else 1 model_kvargs = { "weight_dir": self.weight_dir, "max_total_token_num": max_total_token_num, @@ -146,7 +145,6 @@ def init_model(self, kvargs): "quant_cfg": kvargs.get("quant_cfg", None), "expert_dtype": kvargs.get("expert_dtype", None), "run_mode": self.run_mode, - "decode_batch_multiplier": target_decode_batch_multiplier, } self.model, self.is_multimodal = get_model(model_cfg, model_kvargs) self.model: TpPartBaseModel = self.model # for easy typing @@ -324,15 +322,6 @@ def init_mtp_draft_model(self, main_kvargs: dict): for i in range(draft_model_count): draft_model_cfg, _ = PretrainedConfig.get_config_dict(draft_model_dirs[i]) - if is_chained_draft or ( - self.enable_decode_microbatch_overlap and spec_mode in ("eagle_with_att", "eagle_no_att") - ): - draft_decode_batch_multiplier = self.max_draft_step + 1 - elif spec_mode in ("dspark", "dflash"): - block_size = int(draft_model_cfg["block_size"]) - draft_decode_batch_multiplier = block_size - else: - draft_decode_batch_multiplier = 1 draft_model_kvargs = { "weight_dir": draft_model_dirs[i], "max_total_token_num": self.model.mem_manager.size, @@ -354,7 +343,6 @@ def init_mtp_draft_model(self, main_kvargs: dict): "run_mode": "normal", "main_model": self.model, "mtp_previous_draft_models": self.draft_models.copy(), - "decode_batch_multiplier": draft_decode_batch_multiplier, } draft_model_class = get_draft_model_class( diff --git a/unit_tests/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py new file mode 100644 index 0000000000..4428366f1b --- /dev/null +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -0,0 +1,47 @@ +from types import SimpleNamespace + +import pytest + +import lightllm.common.basemodel.mtp_manager as mtp_manager_module +from lightllm.common.basemodel.mtp_manager import MtpManager + + +@pytest.fixture(autouse=True) +def _reset_mtp_manager(): + MtpManager._instance = None + yield + MtpManager._instance = None + + +def _decode_batch_multiplier(monkeypatch, spec_mode, *, is_draft_model, mtp_step=7): + args = SimpleNamespace( + mtp_mode=spec_mode, + mtp_step=mtp_step, + ) + monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) + return MtpManager.get_instance().get_decode_batch_multiplier(is_draft_model) + + +@pytest.mark.parametrize( + "spec_mode,is_draft_model,expected", + [ + (None, False, 1), + ("eagle3", False, 8), + ("eagle3", True, 1), + ("vanilla_with_att", True, 8), + ("vanilla_no_att", True, 8), + ("eagle_with_att", True, 1), + ("eagle_no_att", True, 1), + ("dspark", True, 7), + ("dflash", True, 7), + ], +) +def test_decode_batch_multiplier(monkeypatch, spec_mode, is_draft_model, expected): + assert _decode_batch_multiplier(monkeypatch, spec_mode, is_draft_model=is_draft_model) == expected + + +def test_get_instance_returns_singleton(monkeypatch): + args = SimpleNamespace(mtp_mode="eagle3", mtp_step=7) + monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) + + assert MtpManager.get_instance() is MtpManager.get_instance() From c65ae9c4199009ca902e7fa1af54457e3bf0955a Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 14 Aug 2026 08:47:01 +0000 Subject: [PATCH 012/103] refactor: clarify MTP CUDA graph batch sizing --- lightllm/common/basemodel/attention/fa3/fp.py | 2 +- .../common/basemodel/attention/fa3/mla.py | 2 +- lightllm/common/basemodel/basemodel.py | 28 ++++++++----- lightllm/common/basemodel/cuda_graph.py | 36 +++++++++------- lightllm/common/basemodel/mtp_manager.py | 14 ++++++- .../basemodel/test_cuda_graph_layout.py | 41 +++++++++++-------- .../common/basemodel/test_mtp_manager.py | 27 ++++++++++-- 7 files changed, 102 insertions(+), 48 deletions(-) diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index 7379957b7e..ccd57a752f 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -26,7 +26,7 @@ def get_page_table_buffer(self): max_att_batch_size = model.graph_max_batch_size if not get_env_start_args().mtp_dynamic_verify: # FA3 merges each fixed speculative block into one attention sequence. - max_att_batch_size //= model.decode_batch_multiplier + max_att_batch_size //= model.mtp_manager.get_decode_batch_multiplier(model.is_mtp_draft_model) buffer_count = 2 if model.args.enable_decode_microbatch_overlap else 1 self._shared_page_table_buffer = [ diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index 2e7f127fbb..65e234abe0 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -24,7 +24,7 @@ def get_page_table_buffer(self): max_att_batch_size = model.graph_max_batch_size if not get_env_start_args().mtp_dynamic_verify: # FA3 merges each fixed speculative block into one attention sequence. - max_att_batch_size //= model.decode_batch_multiplier + max_att_batch_size //= model.mtp_manager.get_decode_batch_multiplier(model.is_mtp_draft_model) buffer_count = 2 if model.args.enable_decode_microbatch_overlap else 1 self._shared_page_table_buffer = [ diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index a2c9553c35..5b9be447a4 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -93,8 +93,9 @@ def __init__(self, kvargs): else self.graph_max_batch_size ) self.mtp_manager = MtpManager.get_instance() - self.decode_batch_multiplier = self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) - self.graph_max_batch_size = self.graph_max_batch_size * self.decode_batch_multiplier + self.graph_max_batch_size = self.graph_max_batch_size * self.mtp_manager.get_decode_batch_multiplier( + self.is_mtp_draft_model + ) self.graph_max_len_in_batch = kvargs.get("graph_max_len_in_batch", 8192) self.disable_cudagraph = kvargs.get("disable_cudagraph", False) @@ -275,17 +276,18 @@ def _init_att_backend1(self): return def _init_cudagraph(self): - batch_multiplier = ( - 1 if self.args.mtp_dynamic_verify and not self.is_mtp_draft_model else self.decode_batch_multiplier - ) + decode_batch_multiplier = self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) + cuda_graph_grow_step_size = self.mtp_manager.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model) self.graph = ( None if self.disable_cudagraph else CudaGraph( + batch_step_size_before_split=cuda_graph_grow_step_size, + split_batch_size=self.args.graph_split_batch_size * decode_batch_multiplier, + batch_step_size_after_split=self.args.graph_grow_step_size * cuda_graph_grow_step_size, max_batch_size=self.graph_max_batch_size, max_len_in_batch=self.graph_max_len_in_batch, tp_world_size=self.tp_world_size_, - batch_multiplier=batch_multiplier, capture_infer_cost=self.args.mtp_dynamic_verify, ) ) @@ -346,15 +348,19 @@ def _full_att_decode_autotune(self): from lightllm.utils.sgl_utils import fa3_decode_autotune - batch_multiplier = ( - 1 if self.args.mtp_dynamic_verify and not self.is_mtp_draft_model else self.decode_batch_multiplier - ) + decode_batch_multiplier = self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) + cuda_graph_grow_step_size = self.mtp_manager.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model) cuda_graph_batch_sizes = CudaGraph.gen_cuda_graph_batch_sizes( + batch_step_size_before_split=cuda_graph_grow_step_size, + split_batch_size=self.args.graph_split_batch_size * decode_batch_multiplier, + batch_step_size_after_split=self.args.graph_grow_step_size * cuda_graph_grow_step_size, max_batch_size=self.graph_max_batch_size, tp_world_size=self.tp_world_size_, - batch_multiplier=batch_multiplier, ) - fa3_decode_autotune(self, cuda_graph_batch_sizes, batch_multiplier=batch_multiplier) + cuda_graph_batch_sizes = [ + batch_size for batch_size in cuda_graph_batch_sizes if batch_size % decode_batch_multiplier == 0 + ] + fa3_decode_autotune(self, cuda_graph_batch_sizes, batch_multiplier=decode_batch_multiplier) return def _init_custom(self): diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 215caffdca..ddf24af02f 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -24,21 +24,25 @@ class CudaGraph: @staticmethod def gen_cuda_graph_batch_sizes( - max_batch_size: int = 8, + batch_step_size_before_split: int, + split_batch_size: int, + batch_step_size_after_split: int, + max_batch_size: int, tp_world_size: int = 1, - batch_multiplier: int = 1, ): args = get_env_start_args() - # gen cuda graph batch_sizes - # cuda graph gen for batch size = [1, 2, 3, ..., graph_split_batch_size] - # and [graph_split_batch_size + graph_grow_step_size, - # if the batch_multiplier is not 1, then the batch_sizes will be multiply of batch_multiplier + # Generate CUDA Graph batch sizes in two phases with independent steps: + # use batch_step_size_before_split up to split_batch_size, then use + # batch_step_size_after_split above it. For example, given + # batch_step_size_before_split=8, split_batch_size=32, + # batch_step_size_after_split=16, and max_batch_size=80, the result is + # [8, 16, 24, 32, 48, 64, 80]. max_batch_size is always included. - split_size = args.graph_split_batch_size * batch_multiplier - grow_size = args.graph_grow_step_size * batch_multiplier - batch_sizes = [i * batch_multiplier for i in range(1, args.graph_split_batch_size + 1)] - batch_sizes.extend(range(split_size + grow_size, max_batch_size, grow_size)) + batch_sizes = list(range(batch_step_size_before_split, split_batch_size + 1, batch_step_size_before_split)) + batch_sizes.extend( + range(split_batch_size + batch_step_size_after_split, max_batch_size, batch_step_size_after_split) + ) batch_sizes = sorted({size for size in batch_sizes if size < max_batch_size} | {max_batch_size}) if args.enable_tpsp_mix_mode: @@ -48,10 +52,12 @@ def gen_cuda_graph_batch_sizes( def __init__( self, + batch_step_size_before_split: int, + split_batch_size: int, + batch_step_size_after_split: int, max_batch_size=8, max_len_in_batch=8192, tp_world_size: int = 1, - batch_multiplier: int = 1, capture_infer_cost: bool = False, ): self.graph = {} @@ -66,9 +72,11 @@ def __init__( self.infer_cost_ms_by_batch_size = {} self.cuda_graph_batch_sizes = self.gen_cuda_graph_batch_sizes( + batch_step_size_before_split=batch_step_size_before_split, + split_batch_size=split_batch_size, + batch_step_size_after_split=batch_step_size_after_split, max_batch_size=self.max_batch_size, tp_world_size=self.tp_world_size, - batch_multiplier=batch_multiplier, ) logger.info(f"cuda graph batch_sizes: {self.cuda_graph_batch_sizes}") @@ -238,7 +246,7 @@ def warmup(self, model): from .basemodel import TpPartBaseModel model: TpPartBaseModel = model - draft_step = model.decode_batch_multiplier - 1 + draft_step = model.mtp_manager.get_decode_batch_multiplier(model.is_mtp_draft_model) - 1 # decode cuda graph init for batch_size in self.cuda_graph_batch_sizes[::-1]: @@ -298,7 +306,7 @@ def warmup_overlap(self, model): from .basemodel import TpPartBaseModel model: TpPartBaseModel = model - draft_step = model.decode_batch_multiplier - 1 + draft_step = model.mtp_manager.get_decode_batch_multiplier(model.is_mtp_draft_model) - 1 for batch_size in self.cuda_graph_batch_sizes[::-1]: decode_batches = [] diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index c914eda891..3450aae1af 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -36,7 +36,7 @@ def get_decode_batch_multiplier(self, is_draft_model: bool) -> int: # Chained MTP runs every draft module over the expanded verify layout. if spec_mode in self._CHAINED_DRAFT_MODES: - return verify_width + return 1 # Recurrent EAGLE draft models decode one row per logical request. if spec_mode in self._RECURRENT_DRAFT_MODES: @@ -47,3 +47,15 @@ def get_decode_batch_multiplier(self, is_draft_model: bool) -> int: return self.args.mtp_step return 1 + + def get_decode_cuda_graph_grow_step_size(self, is_draft_model: bool) -> int: + """Return the batch-size stride used to capture decode CUDA Graphs.""" + + # Draft model CUDA Graphs follow the drafter's physical decode layout. + if is_draft_model: + return self.get_decode_batch_multiplier(is_draft_model=True) + # Main model CUDA Graphs use unit growth for dynamically compacted verify rows. + else: + if self.args.mtp_dynamic_verify: + return 1 + return self.get_decode_batch_multiplier(is_draft_model=False) diff --git a/unit_tests/common/basemodel/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py index 1ba29894ef..17ddcd98da 100644 --- a/unit_tests/common/basemodel/test_cuda_graph_layout.py +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -9,8 +9,6 @@ @pytest.fixture(autouse=True) def _graph_args(monkeypatch): args = SimpleNamespace( - graph_split_batch_size=4, - graph_grow_step_size=2, enable_decode_microbatch_overlap=False, enable_tpsp_mix_mode=False, enable_torch_memory_saver=False, @@ -19,11 +17,13 @@ def _graph_args(monkeypatch): return args -def _batch_sizes(max_batch_size, batch_multiplier=1): - physical_max_batch_size = max_batch_size * batch_multiplier +def _batch_sizes(max_batch_size, batch_stride=1): + physical_max_batch_size = max_batch_size * batch_stride graph = CudaGraph( + batch_step_size_before_split=batch_stride, + split_batch_size=4 * batch_stride, + batch_step_size_after_split=2 * batch_stride, max_batch_size=physical_max_batch_size, - batch_multiplier=batch_multiplier, ) return graph.cuda_graph_batch_sizes @@ -33,7 +33,12 @@ def test_dynamic_schedule_uses_compacted_physical_rows(_graph_args): def test_public_static_schedule_preserves_original_static_mtp_default(_graph_args): - assert CudaGraph.gen_cuda_graph_batch_sizes(max_batch_size=32, batch_multiplier=8) == [ + assert CudaGraph.gen_cuda_graph_batch_sizes( + batch_step_size_before_split=8, + split_batch_size=32, + batch_step_size_after_split=16, + max_batch_size=32, + ) == [ 8, 16, 24, @@ -42,26 +47,28 @@ def test_public_static_schedule_preserves_original_static_mtp_default(_graph_arg def test_instance_and_public_static_schedule_match(_graph_args): - graph = CudaGraph(max_batch_size=128, batch_multiplier=8) + graph = CudaGraph( + batch_step_size_before_split=8, + split_batch_size=32, + batch_step_size_after_split=16, + max_batch_size=128, + ) assert graph.cuda_graph_batch_sizes == CudaGraph.gen_cuda_graph_batch_sizes( + batch_step_size_before_split=8, + split_batch_size=32, + batch_step_size_after_split=16, max_batch_size=graph.max_batch_size, tp_world_size=graph.tp_world_size, - batch_multiplier=8, ) -def test_legacy_vanilla_layout_can_keep_k_plus_one_stride(_graph_args): - assert _batch_sizes(max_batch_size=4, batch_multiplier=8) == [ - 8, - 16, - 24, - 32, - ] +def test_batch_step_size_before_split_controls_capture_range(_graph_args): + assert _batch_sizes(max_batch_size=4, batch_stride=8) == [8, 16, 24, 32] -def test_block_draft_layout_only_pads_to_complete_physical_blocks(_graph_args): - assert _batch_sizes(max_batch_size=8, batch_multiplier=7) == [ +def test_batch_step_size_after_split_controls_capture_range(_graph_args): + assert _batch_sizes(max_batch_size=8, batch_stride=7) == [ 7, 14, 21, diff --git a/unit_tests/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py index 4428366f1b..f2c810771a 100644 --- a/unit_tests/common/basemodel/test_mtp_manager.py +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -17,6 +17,7 @@ def _decode_batch_multiplier(monkeypatch, spec_mode, *, is_draft_model, mtp_step args = SimpleNamespace( mtp_mode=spec_mode, mtp_step=mtp_step, + mtp_dynamic_verify=False, ) monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) return MtpManager.get_instance().get_decode_batch_multiplier(is_draft_model) @@ -28,8 +29,8 @@ def _decode_batch_multiplier(monkeypatch, spec_mode, *, is_draft_model, mtp_step (None, False, 1), ("eagle3", False, 8), ("eagle3", True, 1), - ("vanilla_with_att", True, 8), - ("vanilla_no_att", True, 8), + ("vanilla_with_att", True, 1), + ("vanilla_no_att", True, 1), ("eagle_with_att", True, 1), ("eagle_no_att", True, 1), ("dspark", True, 7), @@ -40,8 +41,28 @@ def test_decode_batch_multiplier(monkeypatch, spec_mode, is_draft_model, expecte assert _decode_batch_multiplier(monkeypatch, spec_mode, is_draft_model=is_draft_model) == expected +@pytest.mark.parametrize( + "dynamic_verify,is_draft_model,expected", + [ + (False, False, 8), + (True, False, 1), + (False, True, 1), + (True, True, 1), + ], +) +def test_decode_cuda_graph_grow_step_size(monkeypatch, dynamic_verify, is_draft_model, expected): + args = SimpleNamespace( + mtp_mode="vanilla_with_att", + mtp_step=7, + mtp_dynamic_verify=dynamic_verify, + ) + monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) + + assert MtpManager.get_instance().get_decode_cuda_graph_grow_step_size(is_draft_model) == expected + + def test_get_instance_returns_singleton(monkeypatch): - args = SimpleNamespace(mtp_mode="eagle3", mtp_step=7) + args = SimpleNamespace(mtp_mode="eagle3", mtp_step=7, mtp_dynamic_verify=False) monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) assert MtpManager.get_instance() is MtpManager.get_instance() From dee474918c1514f196eb27736381ecbc222c5877 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 14 Aug 2026 09:04:56 +0000 Subject: [PATCH 013/103] refactor: centralize MTP decode draft step --- .../common/basemodel/attention/base_att.py | 5 +++-- lightllm/common/basemodel/attention/fa3/fp.py | 4 ++-- lightllm/common/basemodel/attention/fa3/mla.py | 4 ++-- .../common/basemodel/attention/linear/gdn.py | 9 +++++---- .../common/basemodel/attention/triton/fp.py | 2 +- lightllm/common/basemodel/basemodel.py | 1 - lightllm/common/basemodel/batch_objs.py | 2 -- lightllm/common/basemodel/cuda_graph.py | 6 ------ lightllm/common/basemodel/infer_struct.py | 2 -- lightllm/common/basemodel/mtp_manager.py | 5 +++++ .../basemodel/triton_kernel/mtp_utils.py | 7 +++++-- .../mode_backend/generic_padded_pre_process.py | 1 - .../mode_backend/generic_pre_process.py | 1 - .../speculative/proposers/eagle_mtp.py | 2 -- .../speculative/proposers/parallel_block.py | 1 - .../common/basemodel/test_mtp_manager.py | 18 ++++++++++++++++++ .../basemodel/triton_kernel/test_mtp_utils.py | 11 +++++++++-- unit_tests/utils/test_speculative_utils.py | 6 +++--- 18 files changed, 53 insertions(+), 34 deletions(-) diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index cad0563f20..00176b973f 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -39,9 +39,10 @@ def create_att_prefill_state(self) -> "BasePrefillAttState": def create_att_decode_state(self) -> "BaseDecodeAttState": raise NotImplementedError("not impl") - def uses_dynamic_spec_verify_layout(self, infer_state: "InferStateInfo") -> bool: + def uses_dynamic_spec_verify_layout(self) -> bool: args = get_env_start_args() - if infer_state.draft_step == 0 or not args.mtp_dynamic_verify: + draft_step = self.model.mtp_manager.get_decode_draft_step(self.model.is_mtp_draft_model) + if draft_step == 0 or not args.mtp_dynamic_verify: return False # Target verification may compact each request to a different row count. diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index ccd57a752f..311e751732 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -133,9 +133,9 @@ class Fa3DecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: Fa3AttBackend = self.backend - draft_step = self.infer_state.draft_step + draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) decode_rows_per_request = draft_step + 1 - uses_dynamic_spec_verify_layout = self.backend.uses_dynamic_spec_verify_layout(self.infer_state) + uses_dynamic_spec_verify_layout = self.backend.uses_dynamic_spec_verify_layout() if draft_step > 0 and not uses_dynamic_spec_verify_layout: assert self.infer_state.batch_size % decode_rows_per_request == 0, ( diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index 65e234abe0..bf32efb7b4 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -113,9 +113,9 @@ class MlaFa3DecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: MlaFa3AttBackend = self.backend - draft_step = self.infer_state.draft_step + draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) decode_rows_per_request = draft_step + 1 - uses_dynamic_spec_verify_layout = self.backend.uses_dynamic_spec_verify_layout(self.infer_state) + uses_dynamic_spec_verify_layout = self.backend.uses_dynamic_spec_verify_layout() if draft_step > 0 and not uses_dynamic_spec_verify_layout: assert self.infer_state.batch_size % decode_rows_per_request == 0, ( diff --git a/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index e1e90d89dd..44cb08815e 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -206,7 +206,7 @@ class LinearAttDecodeAttState(BaseDecodeAttState): b_num_accepted_tokens: torch.Tensor = None def init_state(self): - draft_step = self.infer_state.draft_step + draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) if draft_step == 0: self.b_conv_buffer_idx = self.infer_state.b_req_idx @@ -215,7 +215,7 @@ def init_state(self): batch_size = self.infer_state.batch_size device = self.infer_state.b_req_idx.device - if self.backend.uses_dynamic_spec_verify_layout(self.infer_state): + if self.backend.uses_dynamic_spec_verify_layout(): ( self.b1_spec_cu_q_seq_len, self.b_conv_buffer_idx, @@ -263,7 +263,8 @@ def decode_att( mixed_qkv, z, b, a = backend._split_qkvzba(mixed_qkvzba) conv_states, ssm_states = self.infer_state.req_manager.get_mamba_cache(layer_num) - if self.infer_state.draft_step > 0: + draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) + if draft_step > 0: core_attn_out = self._gdn_spec_kernel( mixed_qkv, conv_states, @@ -355,7 +356,7 @@ def _gdn_spec_kernel( mixed_qkv, conv_states, layer_weight.linear_conv1d.mm_param.weight, - mtp_step=infer_state.draft_step, + mtp_step=backend.model.mtp_manager.get_decode_draft_step(backend.model.is_mtp_draft_model), bias=layer_weight.linear_conv1d.bias, activation=backend.activation, conv_state_indices=self.b_conv_buffer_idx, diff --git a/lightllm/common/basemodel/attention/triton/fp.py b/lightllm/common/basemodel/attention/triton/fp.py index f2c1127fbb..9a57f2dc28 100644 --- a/lightllm/common/basemodel/attention/triton/fp.py +++ b/lightllm/common/basemodel/attention/triton/fp.py @@ -112,7 +112,7 @@ def decode_att( assert att_control.tp_alibi is not None return self._alibi_decode_att(q=q, k=k, v=v, att_control=att_control, alloc_func=alloc_func) else: - draft_step = self.infer_state.draft_step + draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) q_head_num = q.shape[1] k_head_num = k.shape[1] diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 5b9be447a4..10d752133d 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -426,7 +426,6 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) # 特殊模型,特殊模式的特定变量初始化操作。 infer_state.mtp_draft_input_hiddens = model_input.mtp_draft_input_hiddens - infer_state.draft_step = model_input.draft_step if infer_state.is_prefill: infer_state.prefill_att_state = self.prefill_att_backend.create_att_prefill_state(infer_state=infer_state) diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 69d3ef369c..cd103ec216 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -53,8 +53,6 @@ class ModelInput: # mtp_draft_input_hiddens 用于模型 mtp 模式下 # 的 draft 模型的输入 mtp_draft_input_hiddens: Optional[torch.Tensor] = None - # Maximum number of extra query rows per request. - draft_step: int = 0 def to_cuda(self): self.check_input() diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index ddf24af02f..58e2f90bb2 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -246,8 +246,6 @@ def warmup(self, model): from .basemodel import TpPartBaseModel model: TpPartBaseModel = model - draft_step = model.mtp_manager.get_decode_batch_multiplier(model.is_mtp_draft_model) - 1 - # decode cuda graph init for batch_size in self.cuda_graph_batch_sizes[::-1]: seq_len = 2 @@ -274,7 +272,6 @@ def warmup(self, model): b_mtp_index=b_mtp_index, b_mark_shared_group=b_mark_shared_group, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), - draft_step=draft_step, is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], **model._gen_special_model_input(batch_size), @@ -306,8 +303,6 @@ def warmup_overlap(self, model): from .basemodel import TpPartBaseModel model: TpPartBaseModel = model - draft_step = model.mtp_manager.get_decode_batch_multiplier(model.is_mtp_draft_model) - 1 - for batch_size in self.cuda_graph_batch_sizes[::-1]: decode_batches = [] for micro_batch_index in [0, 1]: @@ -337,7 +332,6 @@ def warmup_overlap(self, model): b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), - draft_step=draft_step, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], **model._gen_special_model_input(batch_size), ) diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index 28d8a5f099..feabe4ac11 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -96,8 +96,6 @@ def __init__(self): # 在开启 mtp_mode 时,mtp draft model # 的输入会用到,其他模型和场景都不会用到 self.mtp_draft_input_hiddens: Optional[torch.Tensor] = None - # Maximum number of extra query rows per request. - self.draft_step: int = 0 # 在单节点多dp的运行模式下,在进行prefill的阶段,如果出现了dp之间数据不平衡的现象, # 可以将推理的数据,进行重新分配到各个dp,在做 att 之前,重新 all to all 到各自的 diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index 3450aae1af..bb0ba57591 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -59,3 +59,8 @@ def get_decode_cuda_graph_grow_step_size(self, is_draft_model: bool) -> int: if self.args.mtp_dynamic_verify: return 1 return self.get_decode_batch_multiplier(is_draft_model=False) + + def get_decode_draft_step(self, is_draft_model: bool) -> int: + """Return the number of extra decode rows processed per request.""" + + return self.get_decode_batch_multiplier(is_draft_model) - 1 diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index cb8a3dc027..90cc4ec4be 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -4,6 +4,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput +from lightllm.common.basemodel.mtp_manager import MtpManager from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import sample_dynamic_spec_row_mask from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size @@ -382,6 +383,7 @@ def _compact_decode_model_input( model_input: ModelInput, selected_row_mask: torch.Tensor, dynamic_batch_size: int, + max_draft_step: int, ) -> ModelInput: assert not model_input.is_prefill assert selected_row_mask.is_cuda @@ -474,7 +476,7 @@ def _compact_decode_model_input( model_input.b_shared_seq_len = out_b_shared_seq_len model_input.b_mark_shared_group = _rebuild_mtp_group_markers( out_b_req_idx, - max_request_rows=model_input.draft_step + 1, + max_request_rows=max_draft_step + 1, ) if model_input.mtp_draft_input_hiddens is not None: @@ -503,7 +505,7 @@ def prepare_dynamic_spec_model_input( assert req_to_next_token_scores is not None assert dynamic_batch_size >= req_num assert dynamic_batch_size <= model_input.batch_size - max_draft_step = int(model_input.draft_step) + max_draft_step = MtpManager.get_instance().get_decode_draft_step(is_draft_model=False) pre_draft_step = max_draft_step if pre_draft_step is None else int(pre_draft_step) assert 0 <= pre_draft_step <= max_draft_step assert model_input.batch_size == req_num * (max_draft_step + 1) @@ -524,6 +526,7 @@ def prepare_dynamic_spec_model_input( model_input=model_input, selected_row_mask=selected_row_mask, dynamic_batch_size=dynamic_batch_size, + max_draft_step=max_draft_step, ) # Keep CPU mem_indexes unfiltered here. Copying selected_row_mask back to # CPU in this hot path synchronizes the overlap stream; the router frees diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py index 8221185572..757288a508 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py @@ -246,7 +246,6 @@ def padded_prepare_decode_inputs( b_position_delta=b_position_delta, b_shared_seq_len=b_shared_seq_len, b_mark_shared_group=b_mark_shared_group, - draft_step=args_mtp_step, is_prefill=False, multimodal_params=batch_multimodal_params, ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 911929ef89..3bfdc8b063 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -161,7 +161,6 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_position_delta=b_position_delta, b_shared_seq_len=b_shared_seq_len, b_mark_shared_group=b_mark_shared_group, - draft_step=get_env_start_args().mtp_step, is_prefill=False, multimodal_params=multimodal_params, ) diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py index a0ba2557bd..8814e07e99 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -178,7 +178,6 @@ def propose_next( ) draft_input.b_mark_shared_group = torch.ones_like(draft_input.b_req_idx) draft_input.b_shared_seq_len = None - draft_input.draft_step = 0 if len(draft_input.multimodal_params) != request_count: empty_multimodal_params = {"images": [], "audios": []} draft_input.multimodal_params = [empty_multimodal_params] * request_count @@ -315,7 +314,6 @@ def propose_next_overlap( ) model_input.b_mark_shared_group = torch.ones_like(model_input.b_req_idx) model_input.b_shared_seq_len = None - model_input.draft_step = 0 if len(model_input.multimodal_params) != model_input.batch_size: empty_multimodal_params = {"images": [], "audios": []} model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size diff --git a/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py b/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py index 3f75e1b326..73c893b4e7 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py +++ b/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py @@ -92,7 +92,6 @@ def build_block_draft_input( draft_input.batch_size = draft_input.total_token_num draft_input.max_q_seq_len = 1 draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size - draft_input.draft_step = block_size - 1 draft_input.b_req_idx = ( main_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() ) diff --git a/unit_tests/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py index f2c810771a..004d119f41 100644 --- a/unit_tests/common/basemodel/test_mtp_manager.py +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -61,6 +61,24 @@ def test_decode_cuda_graph_grow_step_size(monkeypatch, dynamic_verify, is_draft_ assert MtpManager.get_instance().get_decode_cuda_graph_grow_step_size(is_draft_model) == expected +@pytest.mark.parametrize( + "spec_mode,is_draft_model,expected", + [ + (None, False, 0), + ("eagle3", False, 7), + ("eagle3", True, 0), + ("vanilla_with_att", True, 0), + ("dspark", True, 6), + ("dflash", True, 6), + ], +) +def test_decode_draft_step(monkeypatch, spec_mode, is_draft_model, expected): + args = SimpleNamespace(mtp_mode=spec_mode, mtp_step=7, mtp_dynamic_verify=False) + monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) + + assert MtpManager.get_instance().get_decode_draft_step(is_draft_model) == expected + + def test_get_instance_returns_singleton(monkeypatch): args = SimpleNamespace(mtp_mode="eagle3", mtp_step=7, mtp_dynamic_verify=False) monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py index c8adbe9825..12cd766d67 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -6,10 +6,18 @@ pytest.skip("requires CUDA", allow_module_level=True) from lightllm.common.basemodel.batch_objs import ModelInput +from lightllm.common.basemodel.mtp_manager import MtpManager from lightllm.common.basemodel.triton_kernel import mtp_utils from lightllm.utils.envs_utils import get_env_start_args +@pytest.fixture(autouse=True) +def _reset_mtp_manager(): + MtpManager._instance = None + yield + MtpManager._instance = None + + def test_compact_dynamic_spec_model_input(monkeypatch): monkeypatch.setenv( "LIGHTLLM_START_ARGS", @@ -43,7 +51,6 @@ def test_compact_dynamic_spec_model_input(monkeypatch): mem_indexes=torch.arange(12, dtype=torch.int32, device="cuda") + 100, mem_indexes_cpu=torch.arange(12, dtype=torch.int32, device="cpu") + 100, is_prefill=False, - draft_step=3, multimodal_params=[{"row": i, "images": [], "audios": []} for i in range(12)], mtp_draft_input_hiddens=(torch.arange(12 * 5, dtype=torch.float32, device="cuda").reshape(12, 5) + 0.5), ) @@ -117,7 +124,6 @@ def test_compaction_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size( b_mark_shared_group=torch.tensor([0, 0, 0, 0, 5], dtype=torch.int32, device="cuda"), mem_indexes=torch.arange(5, dtype=torch.int32, device="cuda"), mem_indexes_cpu=torch.arange(5, dtype=torch.int32, device="cpu"), - draft_step=4, is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(5)], ) @@ -127,6 +133,7 @@ def test_compaction_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size( model_input=model_input, selected_row_mask=selected_row_mask, dynamic_batch_size=5, + max_draft_step=4, ) torch.cuda.synchronize() diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 979396f6a5..24612cf600 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -53,7 +53,7 @@ def test_qwen3_eagle_uses_layers_checkpoint_prefix(): ("dspark", True, 7, True, False), ("dflash", False, 7, True, True), ("dflash", True, 7, True, False), - ("vanilla_with_att", True, 7, True, True), + ("vanilla_with_att", True, 0, True, False), ("eagle3", True, 0, True, False), ("eagle3", False, 7, False, False), ], @@ -74,11 +74,11 @@ def test_attention_backend_selects_dynamic_spec_layout( backend = SimpleNamespace( model=SimpleNamespace( is_mtp_draft_model=is_draft_model, + mtp_manager=SimpleNamespace(get_decode_draft_step=lambda _: draft_step), ) ) - infer_state = SimpleNamespace(draft_step=draft_step) - assert BaseAttBackend.uses_dynamic_spec_verify_layout(backend, infer_state) is expected + assert BaseAttBackend.uses_dynamic_spec_verify_layout(backend) is expected @pytest.mark.parametrize( From fe70b2ff853352d09951ddd0102d514efb1df97d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 06:05:00 +0000 Subject: [PATCH 014/103] refactor: isolate hidden collector inference state --- lightllm/common/basemodel/basemodel.py | 79 +++--- lightllm/common/basemodel/hidden_collector.py | 227 +++++++++++------- lightllm/common/basemodel/infer_struct.py | 4 + lightllm/common/basemodel/mtp_manager.py | 28 ++- .../common/basemodel/prefill_cuda_graph.py | 26 +- .../common/basemodel/test_hidden_collector.py | 167 +++++++------ .../common/basemodel/test_mtp_manager.py | 37 +++ .../test_prefill_cuda_graph_state.py | 83 +++++++ unit_tests/utils/test_speculative_utils.py | 16 +- 9 files changed, 455 insertions(+), 212 deletions(-) create mode 100644 unit_tests/common/basemodel/test_prefill_cuda_graph_state.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 10d752133d..d4a036502b 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -29,7 +29,7 @@ from lightllm.distributed.communication_op import dist_group_manager from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.hidden_collector import ( - HiddenCollector, + NoopHiddenCollector, unpad_collected_hidden, ) from lightllm.common.basemodel.mtp_manager import MtpManager @@ -367,14 +367,7 @@ def _init_custom(self): pass def _init_hidden_collector(self): - microbatch_count = ( - 2 if self.args.enable_prefill_microbatch_overlap or self.args.enable_decode_microbatch_overlap else 1 - ) - self.hidden_collector = HiddenCollector( - model=self, - spec_mode=self.args.mtp_mode, - microbatch_count=microbatch_count, - ) + self.hidden_collector_prototype = self.mtp_manager.create_hidden_collector(model=self) @torch.no_grad() def forward(self, model_input: ModelInput): @@ -388,6 +381,7 @@ def forward(self, model_input: ModelInput): def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() + infer_state.hidden_collector = self.hidden_collector_prototype.new_instance() infer_state.input_ids = model_input.input_ids infer_state.is_prefill = model_input.is_prefill infer_state.is_token_healing = self.is_token_healing @@ -698,19 +692,21 @@ def _context_forward(self, infer_state: InferStateInfo): input_embs = self.pre_infer._tpsp_sp_split(input=input_embs, infer_state=infer_state) input_tensors = [input_embs] - hidden_collector = HiddenCollector() if Autotuner.is_autotune_warmup() else self.hidden_collector + if Autotuner.is_autotune_warmup(): + infer_state.hidden_collector = NoopHiddenCollector() - def prefill_func(input_tensors, infer_state): + def prefill_func(input_tensors, _infer_state): + hidden_collector = _infer_state.hidden_collector _input_embs = input_tensors[0] for i in range(self.layers_num): layer = self.layers_infer[i] - _input_embs = layer.context_forward(_input_embs, infer_state, self.trans_layers_weight[i]) + _input_embs = layer.context_forward(_input_embs, _infer_state, self.trans_layers_weight[i]) hidden_collector.add( layer_index=i, hidden=_input_embs, ) - return hidden_collector.prefill_outputs(_input_embs) + return [_input_embs] handle_token_num = infer_state.input_ids.shape[0] @@ -742,11 +738,9 @@ def prefill_func(input_tensors, infer_state): last_input_embs = infer_state._all_to_all_unbalance_get(data=last_input_embs) predict_logits = self.post_infer.token_forward(last_input_embs, infer_state, self.pre_post_weight) - spec_hidden = hidden_collector.finish( - infer_state=infer_state, - final_hidden=last_input_embs, - forward_outputs=output_tensors, - ) + hidden_collector = infer_state.hidden_collector + hidden_collector.add_final_hidden(last_input_embs) + spec_hidden = hidden_collector.finish(infer_state=infer_state) model_output = ModelOutput( logits=predict_logits, spec_hidden=spec_hidden, @@ -759,6 +753,7 @@ def prefill_func(input_tensors, infer_state): return model_output def _token_forward(self, infer_state: InferStateInfo): + hidden_collector = infer_state.hidden_collector input_ids = infer_state.input_ids cuda_input_ids = input_ids input_embs = self.pre_infer.token_forward(cuda_input_ids, infer_state, self.pre_post_weight) @@ -767,17 +762,15 @@ def _token_forward(self, infer_state: InferStateInfo): for i in range(self.layers_num): layer = self.layers_infer[i] input_embs: torch.Tensor = layer.token_forward(input_embs, infer_state, self.trans_layers_weight[i]) - self.hidden_collector.add(layer_index=i, hidden=input_embs) + hidden_collector.add(layer_index=i, hidden=input_embs) last_input_embs = self.post_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) predict_logits: torch.Tensor = self.post_infer.token_forward( last_input_embs, infer_state=infer_state, layer_weight=self.pre_post_weight ) - spec_hidden = self.hidden_collector.finish( - infer_state=infer_state, - final_hidden=last_input_embs, - ) + hidden_collector.add_final_hidden(last_input_embs) + spec_hidden = hidden_collector.finish(infer_state=infer_state) model_output = ModelOutput(logits=predict_logits.contiguous(), spec_hidden=spec_hidden) # 在 cuda graph 模式下,输出需要转为 no ref tensor, 加强mem pool 的复用,降低显存的使用。 @@ -980,6 +973,8 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode @final def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state1: InferStateInfo): g_cache_manager.cache_env_in() + hidden_collector0 = infer_state.hidden_collector + hidden_collector1 = infer_state1.hidden_collector input_embs, input_embs1 = self.pre_infer.overlap_tpsp_context_forward( infer_state.input_ids, infer_state1.input_ids, infer_state, infer_state1, self.pre_post_weight @@ -1000,14 +995,13 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state input_embs, input_embs1 = self.layers_infer[i].overlap_tpsp_context_forward( input_embs, input_embs1, infer_state, infer_state1, self.trans_layers_weight[i] ) - self.hidden_collector.add( + hidden_collector0.add( layer_index=i, hidden=input_embs, ) - self.hidden_collector.add( + hidden_collector1.add( layer_index=i, hidden=input_embs1, - microbatch_index=1, ) # 折叠模式调用完infer_state 和 infer_state1 上的hook函数后,input_embs 和 input_embs1 才具备正确的运算数据。 @@ -1025,15 +1019,10 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state ) g_cache_manager.cache_env_out() - spec_hidden = self.hidden_collector.finish( - infer_state=infer_state, - final_hidden=last_input_embs, - ) - spec_hidden1 = self.hidden_collector.finish( - infer_state=infer_state1, - final_hidden=last_input_embs1, - microbatch_index=1, - ) + hidden_collector0.add_final_hidden(last_input_embs) + hidden_collector1.add_final_hidden(last_input_embs1) + spec_hidden = hidden_collector0.finish(infer_state=infer_state) + spec_hidden1 = hidden_collector1.finish(infer_state=infer_state1) model_output = ModelOutput( logits=predict_logits.contiguous(), spec_hidden=spec_hidden, @@ -1049,6 +1038,8 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state @final def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: InferStateInfo): + hidden_collector0 = infer_state.hidden_collector + hidden_collector1 = infer_state1.hidden_collector input_embs, input_embs1 = self.pre_infer.overlap_tpsp_token_forward( infer_state.input_ids, infer_state1.input_ids, infer_state, infer_state1, self.pre_post_weight ) @@ -1059,14 +1050,13 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: input_embs, input_embs1 = self.layers_infer[i].overlap_tpsp_token_forward( input_embs, input_embs1, infer_state, infer_state1, self.trans_layers_weight[i] ) - self.hidden_collector.add( + hidden_collector0.add( layer_index=i, hidden=input_embs, ) - self.hidden_collector.add( + hidden_collector1.add( layer_index=i, hidden=input_embs1, - microbatch_index=1, ) # 折叠模式调用完infer_state 上的hook函数后,input_embs 和 input_embs 才具备正确的运算数据。 @@ -1080,15 +1070,10 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: last_input_embs, last_input_embs1, infer_state, infer_state1, self.pre_post_weight ) - spec_hidden = self.hidden_collector.finish( - infer_state=infer_state, - final_hidden=last_input_embs, - ) - spec_hidden1 = self.hidden_collector.finish( - infer_state=infer_state1, - final_hidden=last_input_embs1, - microbatch_index=1, - ) + hidden_collector0.add_final_hidden(last_input_embs) + hidden_collector1.add_final_hidden(last_input_embs1) + spec_hidden = hidden_collector0.finish(infer_state=infer_state) + spec_hidden1 = hidden_collector1.finish(infer_state=infer_state1) model_output = ModelOutput(logits=predict_logits.contiguous(), spec_hidden=spec_hidden) model_output1 = ModelOutput(logits=predict_logits1.contiguous(), spec_hidden=spec_hidden1) diff --git a/lightllm/common/basemodel/hidden_collector.py b/lightllm/common/basemodel/hidden_collector.py index 64d67fe69f..c8da1b0b08 100644 --- a/lightllm/common/basemodel/hidden_collector.py +++ b/lightllm/common/basemodel/hidden_collector.py @@ -1,66 +1,174 @@ from __future__ import annotations -from typing import Iterable, List, Optional +import copy +from abc import ABC, abstractmethod +from typing import List, Optional import torch from transformers.configuration_utils import PretrainedConfig from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor def unpad_collected_hidden(hidden: Optional[torch.Tensor], token_count: int) -> Optional[torch.Tensor]: return None if hidden is None else hidden[:token_count] -class NoopHiddenCollector: - """Null object used by models that do not expose speculative features.""" +class HiddenCollector(ABC): + """Hidden state 收集器的抽象基类。 + + 推理过程中,模型会在每一层计算完成后调用 :meth:`add`,并在一次 forward + 结束时调用 :meth:`finish` 生成供投机解码使用的 ``spec_hidden``。不同投机 + 解码模式可以通过子类决定不收集 hidden、只返回最终层 hidden,或者收集若干 + 中间层 hidden。 + + BaseModel 持有一个不承载推理状态的 prototype,每个 InferStateInfo 通过 + :meth:`new_instance` 获得独立收集器,使不同请求和 microbatch 的临时 hidden + state 自然隔离,具体实现无需感知 microbatch 编号。 + """ + + @abstractmethod + def new_instance(self) -> "HiddenCollector": + """基于当前 prototype 创建一个不包含运行时 tensor 的新实例。 + + 新实例可以复用 model、目标层编号等只读配置,但不得共享 ``layer_hiddens`` + 或 ``final_hidden`` 等运行时状态。该接口用于替代 ``deepcopy``,避免递归 + 复制 model、CUDA tensor 和通信对象。 + + Returns: + 与当前 collector 类型和配置一致、运行时状态为空的新实例。 + """ + raise NotImplementedError + + def restore_graph_state(self, graph_collector: "HiddenCollector") -> None: + """从 Prefill CUDA Graph 的 capture collector 恢复本次 replay 所需状态。 + + 基类默认不恢复任何数据。需要中间层 hidden 的子类应只复制保存 tensor + 引用的容器,不复制 tensor 本身,使本次 infer state 可以读取 graph replay + 更新后的固定地址,同时在 :meth:`finish` 后独立清理自己的容器。 + + Args: + graph_collector: capture 阶段保存在 Prefill CUDA Graph 中的只读 collector。 + """ + return + + def release_graph_tensor_ownership(self) -> tuple[int, int]: + """将 Prefill CUDA Graph capture 状态转换为不持有显存所有权的引用。 + + 普通 collector 不应调用该接口。完成 CUDA Graph capture 后,graph memory + pool 会负责保证固定显存地址的生命周期;collector 只需保存指针视图供 + replay 后读取。基类没有需要转换的 tensor,因此返回零统计值。 + + Returns: + ``(tensor_count, total_nbytes)``,分别表示转换的 tensor 数量及其总字节数。 + """ + return 0, 0 def add(self, layer_index: int, hidden: torch.Tensor) -> None: + """接收一个 decoder layer 刚计算完成的本地 hidden state。 + + BaseModel 会按照模型层顺序逐层调用该接口。基类默认不保存 tensor,适用 + 于不需要中间层 hidden 的实现;需要收集中间层的子类应重写该方法,并在 + 必要时 clone 会被后续层复用的输入缓冲区。 + + Args: + layer_index: 当前 decoder layer 的零基索引。 + hidden: 当前层输出的本地 hidden tensor,可能仍是 TP/SP 切分状态。 + """ + return + + def add_final_hidden(self, final_hidden: torch.Tensor) -> None: + """接收完成模型输出侧 gather 后的最终层 hidden tensor。 + + BaseModel 在 logits 计算完成后、调用 :meth:`finish` 前调用该接口。基类 + 默认不保存 tensor;需要直接返回最终层 hidden 的子类应重写该方法并保存 + 引用,随后在 :meth:`finish` 中消费和清理。 + + Args: + final_hidden: 已完成 TP/SP all-gather 和 DP unbalance 的最终层 hidden。 + """ return - def prefill_outputs(self, final_hidden: torch.Tensor) -> List[torch.Tensor]: - return [final_hidden] + @abstractmethod + def finish(self, infer_state) -> Optional[torch.Tensor]: + """结束当前 microbatch 的收集并生成投机解码所需的 hidden tensor。 + + 子类应在此完成必要的拼接、TP/SP all-gather、DP unbalance 和 contiguous + 转换,并在返回前清理当前实例的临时状态。该接口在每次 forward 结束时 + 调用一次;不需要向 drafter 提供 hidden 的实现应返回 ``None``。 + + Args: + infer_state: 当前 forward 的推理状态,包含通信拓扑、DP balance 等信息。 + + Returns: + 提供给投机解码 drafter 的连续 hidden tensor;当前模式不需要 hidden 时 + 返回 ``None``。 + """ + raise NotImplementedError + + +class NoopHiddenCollector(HiddenCollector): + """Null object used by models that do not expose speculative features.""" + + def new_instance(self) -> HiddenCollector: + return NoopHiddenCollector() - def finish( - self, - infer_state, - final_hidden: torch.Tensor, - forward_outputs: Optional[List[torch.Tensor]] = None, - ) -> Optional[torch.Tensor]: + def finish(self, infer_state) -> Optional[torch.Tensor]: return None -class FinalHiddenCollector(NoopHiddenCollector): +class FinalHiddenCollector(HiddenCollector): """Returns the final decoder hidden state without per-layer bookkeeping.""" - def finish( - self, - infer_state, - final_hidden: torch.Tensor, - forward_outputs: Optional[List[torch.Tensor]] = None, - ) -> torch.Tensor: + def __init__(self) -> None: + self.final_hidden: Optional[torch.Tensor] = None + + def new_instance(self) -> HiddenCollector: + return FinalHiddenCollector() + + def add_final_hidden(self, final_hidden: torch.Tensor) -> None: + self.final_hidden = final_hidden + + def finish(self, infer_state) -> torch.Tensor: + assert self.final_hidden is not None + final_hidden = self.final_hidden + self.final_hidden = None return final_hidden.contiguous() -class LayerHiddenCollector(NoopHiddenCollector): +class LayerHiddenCollector(HiddenCollector): """Collects selected decoder-layer outputs for an intermediate-hidden draft.""" - def __init__(self, model, layer_ids: Optional[Iterable[int]] = None) -> None: + def __init__(self, model) -> None: self.model = model self.layer_num = model.layers_num - self.layer_ids = self._resolve_layer_ids(layer_ids) + self.layer_ids = self._load_layer_ids() self.layer_hiddens: List[torch.Tensor] = [] - def _resolve_layer_ids(self, layer_ids: Optional[Iterable[int]]) -> frozenset[int]: + def new_instance(self) -> HiddenCollector: + collector = copy.copy(self) + collector.layer_hiddens = [] + return collector + + def restore_graph_state(self, graph_collector: HiddenCollector) -> None: + assert isinstance(graph_collector, LayerHiddenCollector) + self.layer_hiddens = graph_collector.layer_hiddens.copy() + + def release_graph_tensor_ownership(self) -> tuple[int, int]: + tensor_count = len(self.layer_hiddens) + total_nbytes = sum(hidden.numel() * hidden.element_size() for hidden in self.layer_hiddens) + self.layer_hiddens = [tensor_to_no_ref_tensor(hidden) for hidden in self.layer_hiddens] + return tensor_count, total_nbytes + + def _load_layer_ids(self) -> frozenset[int]: + draft_model_dirs = get_env_start_args().mtp_draft_model_dir + assert draft_model_dirs + draft_config, _ = PretrainedConfig.get_config_dict(draft_model_dirs[0]) + layer_ids = draft_config.get("target_layer_ids") if layer_ids is None: - draft_model_dirs = get_env_start_args().mtp_draft_model_dir - assert draft_model_dirs - draft_config, _ = PretrainedConfig.get_config_dict(draft_model_dirs[0]) - layer_ids = draft_config.get("target_layer_ids") - if layer_ids is None: - layer_ids = draft_config.get("dflash_config", {}).get("target_layer_ids") - if layer_ids is None: - layer_ids = [1, self.layer_num // 2 - 1, self.layer_num - 4] + layer_ids = draft_config.get("dflash_config", {}).get("target_layer_ids") + assert layer_ids is not None, f"target_layer_ids is required in draft config: {draft_model_dirs[0]}" resolved_layer_ids = frozenset(int(layer_id) for layer_id in layer_ids) assert resolved_layer_ids and all( @@ -83,65 +191,10 @@ def _local_hidden(self) -> torch.Tensor: return self.layer_hiddens[0] return torch.cat(self.layer_hiddens, dim=-1) - def prefill_outputs(self, final_hidden: torch.Tensor) -> List[torch.Tensor]: - return [final_hidden, self._local_hidden()] - - def finish( - self, - infer_state, - final_hidden: torch.Tensor, - forward_outputs: Optional[List[torch.Tensor]] = None, - ) -> torch.Tensor: - local_hidden = self._local_hidden() if forward_outputs is None else forward_outputs[1] + def finish(self, infer_state) -> torch.Tensor: + local_hidden = self._local_hidden() self.layer_hiddens.clear() hidden = self.model.pre_infer._tpsp_allgather(input=local_hidden, infer_state=infer_state) if infer_state.need_dp_prefill_balance: hidden = infer_state._all_to_all_unbalance_get(data=hidden) return hidden.contiguous() - - -class HiddenCollector: - """Collect hidden states for one or more independently executed microbatches.""" - - def __init__( - self, - model=None, - spec_mode: Optional[str] = None, - layer_ids: Optional[Iterable[int]] = None, - microbatch_count: int = 1, - ) -> None: - assert microbatch_count > 0 - if spec_mode is not None: - assert model is not None - - collector_kwargs = {} - if spec_mode is None: - collector_type = NoopHiddenCollector - elif model.is_mtp_draft_model: - collector_type = NoopHiddenCollector if spec_mode in ("dspark", "dflash") else FinalHiddenCollector - elif spec_mode not in ("eagle3", "dspark", "dflash"): - collector_type = FinalHiddenCollector - else: - collector_type = LayerHiddenCollector - collector_kwargs = {"model": model, "layer_ids": layer_ids} - - self.collectors = tuple(collector_type(**collector_kwargs) for _ in range(microbatch_count)) - - def add(self, layer_index: int, hidden: torch.Tensor, microbatch_index: int = 0) -> None: - self.collectors[microbatch_index].add(layer_index=layer_index, hidden=hidden) - - def finish( - self, - infer_state, - final_hidden: torch.Tensor, - forward_outputs: Optional[List[torch.Tensor]] = None, - microbatch_index: int = 0, - ) -> Optional[torch.Tensor]: - return self.collectors[microbatch_index].finish( - infer_state=infer_state, - final_hidden=final_hidden, - forward_outputs=forward_outputs, - ) - - def prefill_outputs(self, final_hidden: torch.Tensor, microbatch_index: int = 0) -> List[torch.Tensor]: - return self.collectors[microbatch_index].prefill_outputs(final_hidden) diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index feabe4ac11..5fb5599469 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -71,6 +71,10 @@ def __init__(self): # 在一些细节场景下需要有该信息区分一些资源的申请和管理。 self.microbatch_index: int = 0 + # 当前 forward 独占的 hidden state 收集器。普通推理使用短生命周期实例; + # Prefill CUDA Graph 使用随 graph infer state 长期保存的实例。 + self.hidden_collector = None + # 衍生使用的管理变量,为了方便扩展接入其他的高性能attention推理算子,在 # inferstate 基类上添加下面的标记变量,用于扩展。 # b 开头的tensor变量其shape为[batch_size,] diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index bb0ba57591..1d39f43e3f 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -1,10 +1,16 @@ from typing import ClassVar, Optional +from lightllm.common.basemodel.hidden_collector import ( + FinalHiddenCollector, + HiddenCollector, + LayerHiddenCollector, + NoopHiddenCollector, +) from lightllm.utils.envs_utils import get_env_start_args class MtpManager: - """Manage the physical decode layout for main and draft models.""" + """Manage MTP layout policy and model-local helper construction.""" _instance: ClassVar[Optional["MtpManager"]] = None _CHAINED_DRAFT_MODES = ("vanilla_with_att", "vanilla_no_att") @@ -64,3 +70,23 @@ def get_decode_draft_step(self, is_draft_model: bool) -> int: """Return the number of extra decode rows processed per request.""" return self.get_decode_batch_multiplier(is_draft_model) - 1 + + def create_hidden_collector( + self, + model, + ) -> HiddenCollector: + """Create a model-local hidden-collector prototype for the configured MTP mode.""" + + spec_mode = self.args.mtp_mode + collector_kwargs = {} + if spec_mode is None: + collector_type = NoopHiddenCollector + elif model.is_mtp_draft_model: + collector_type = NoopHiddenCollector if spec_mode in self._BLOCK_DRAFT_MODES else FinalHiddenCollector + elif spec_mode in ("eagle3", *self._BLOCK_DRAFT_MODES): + collector_type = LayerHiddenCollector + collector_kwargs.update(model=model) + else: + collector_type = FinalHiddenCollector + + return collector_type(**collector_kwargs) diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index 1c1148a55d..6f8d91471d 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -89,7 +89,22 @@ def _capture_prefill( graph_input_tensors = [tensor_to_no_ref_tensor(e) for e in graph_input_tensors] graph_out_tensors = [tensor_to_no_ref_tensor(e) for e in graph_out_tensors] - self.graph[handle_token_num] = (infer_state, graph_input_tensors, graph_out_tensors) + graph_hidden_collector = infer_state.hidden_collector + hidden_tensor_count, hidden_tensor_nbytes = graph_hidden_collector.release_graph_tensor_ownership() + logger.info( + f"Prefill CUDA Graph hidden collector 已完成 capture 状态托管:" + f"handle_token_num={handle_token_num}, " + f"collector={graph_hidden_collector.__class__.__name__}, " + f"no_ref_tensor_count={hidden_tensor_count}, " + f"no_ref_tensor_nbytes={hidden_tensor_nbytes}。" + f"这些 tensor 的固定地址由 graph memory pool 管理,collector 仅保留无所有权引用供 replay 使用。" + ) + self.graph[handle_token_num] = ( + infer_state, + graph_input_tensors, + graph_out_tensors, + graph_hidden_collector, + ) self.replay(input_tensors, infer_state) return graph_out_tensors @@ -133,12 +148,19 @@ def capture_prefill( def _replay(self, input_tensors: List[torch.Tensor], infer_state: InferStateInfo) -> List[torch.Tensor]: handle_token_num = infer_state.total_token_num - infer_state.prefix_total_token_num - graph_infer_state, graph_input_tensors, graph_output_tensors = self.graph[handle_token_num] + graph_infer_state, graph_input_tensors, graph_output_tensors, graph_hidden_collector = self.graph[ + handle_token_num + ] graph_infer_state: InferStateInfo = graph_infer_state for graph_in_tensor, in_tensor in zip(graph_input_tensors, input_tensors): graph_in_tensor.copy_(in_tensor) graph_infer_state.copy_for_prefill_cuda_graph(new_infer_state=infer_state) + # 首次 capture 后 replay 时,infer_state 与 graph_infer_state 是同一对象, + # 需要先创建运行时实例,避免 finish 清空 graph 中保存的 capture collector。 + if infer_state.hidden_collector is graph_hidden_collector: + infer_state.hidden_collector = graph_hidden_collector.new_instance() + infer_state.hidden_collector.restore_graph_state(graph_hidden_collector) graph_infer_state.prefill_replay(infer_state) return graph_output_tensors diff --git a/unit_tests/common/basemodel/test_hidden_collector.py b/unit_tests/common/basemodel/test_hidden_collector.py index a352d15bcf..3ce9d881ca 100644 --- a/unit_tests/common/basemodel/test_hidden_collector.py +++ b/unit_tests/common/basemodel/test_hidden_collector.py @@ -1,7 +1,10 @@ +from inspect import isabstract from types import SimpleNamespace +import pytest import torch +import lightllm.common.basemodel.hidden_collector as hidden_collector_module from lightllm.common.basemodel.hidden_collector import ( FinalHiddenCollector, HiddenCollector, @@ -17,133 +20,155 @@ def _tpsp_allgather(input, infer_state): return input -def test_hidden_collector_selects_implementation(): - model = SimpleNamespace(is_mtp_draft_model=False, layers_num=3, pre_infer=_IdentityPreInfer()) - noop_collector = HiddenCollector() - final_hidden_collector = HiddenCollector(model=model, spec_mode="vanilla_with_att") - layer_hidden_collector = HiddenCollector( - model=model, - spec_mode="eagle3", - layer_ids=[0, 2], - microbatch_count=2, +def _mock_target_layer_ids(monkeypatch, layer_ids): + monkeypatch.setattr( + hidden_collector_module, + "get_env_start_args", + lambda: SimpleNamespace(mtp_draft_model_dir=["/models/draft"]), + ) + monkeypatch.setattr( + hidden_collector_module.PretrainedConfig, + "get_config_dict", + lambda _: ({"target_layer_ids": layer_ids}, {}), ) - assert isinstance(noop_collector, HiddenCollector) - assert isinstance(noop_collector.collectors[0], NoopHiddenCollector) - assert isinstance(final_hidden_collector.collectors[0], FinalHiddenCollector) - assert all(isinstance(collector, LayerHiddenCollector) for collector in layer_hidden_collector.collectors) - - -def test_draft_hidden_collector_follows_spec_mode(): - model = SimpleNamespace(is_mtp_draft_model=True) - - autoregressive_collector = HiddenCollector(model=model, spec_mode="eagle3") - parallel_block_collector = HiddenCollector(model=model, spec_mode="dspark") - assert isinstance(autoregressive_collector.collectors[0], FinalHiddenCollector) - assert isinstance(parallel_block_collector.collectors[0], NoopHiddenCollector) +def test_hidden_collector_is_base_class_for_implementations(): + assert isabstract(HiddenCollector) + assert issubclass(NoopHiddenCollector, HiddenCollector) + assert issubclass(FinalHiddenCollector, HiddenCollector) + assert issubclass(LayerHiddenCollector, HiddenCollector) -def test_hidden_collector_supports_single_and_overlap_forward(): - model = SimpleNamespace(is_mtp_draft_model=False) - collector = HiddenCollector(model=model, spec_mode="vanilla_with_att", microbatch_count=2) +def test_final_hidden_collectors_are_independent_instances(): + prototype = FinalHiddenCollector() + collector0 = prototype.new_instance() + collector1 = prototype.new_instance() hidden0 = torch.randn(2, 3) hidden1 = torch.randn(2, 3) infer_state = SimpleNamespace(need_dp_prefill_balance=False) - collector.add(layer_index=0, hidden=hidden0) - collected = collector.finish( - infer_state=infer_state, - final_hidden=hidden0, - ) + collector0.add_final_hidden(hidden0) + collected = collector0.finish(infer_state=infer_state) assert collected.data_ptr() == hidden0.data_ptr() + assert collector0.final_hidden is None - collector.add(layer_index=0, hidden=hidden0) - collector.add(layer_index=0, hidden=hidden1, microbatch_index=1) - collected0 = collector.finish( - infer_state=infer_state, - final_hidden=hidden0, - ) - collected1 = collector.finish( - infer_state=infer_state, - final_hidden=hidden1, - microbatch_index=1, - ) + collector0.add_final_hidden(hidden0) + collector1.add_final_hidden(hidden1) + collected0 = collector0.finish(infer_state=infer_state) + collected1 = collector1.finish(infer_state=infer_state) assert collected0.data_ptr() == hidden0.data_ptr() assert collected1.data_ptr() == hidden1.data_ptr() -def test_layer_hidden_collector_keeps_microbatch_state_separate(): +def test_layer_hidden_collector_keeps_microbatch_state_separate(monkeypatch): + _mock_target_layer_ids(monkeypatch, [0]) model = SimpleNamespace(is_mtp_draft_model=False, layers_num=2, pre_infer=_IdentityPreInfer()) - collector = HiddenCollector(model=model, spec_mode="eagle3", layer_ids=[0], microbatch_count=2) + prototype = LayerHiddenCollector(model=model) + collector0 = prototype.new_instance() + collector1 = prototype.new_instance() hidden0 = torch.full((2, 3), 1.0) hidden1 = torch.full((2, 3), 2.0) infer_state = SimpleNamespace(need_dp_prefill_balance=False) - collector.add(layer_index=0, hidden=hidden0) - collector.add(layer_index=0, hidden=hidden1, microbatch_index=1) + collector0.add(layer_index=0, hidden=hidden0) + collector1.add(layer_index=0, hidden=hidden1) - collected0 = collector.finish(infer_state=infer_state, final_hidden=hidden0) - collected1 = collector.finish(infer_state=infer_state, final_hidden=hidden1, microbatch_index=1) + collected0 = collector0.finish(infer_state=infer_state) + collected1 = collector1.finish(infer_state=infer_state) assert torch.equal(collected0, hidden0) assert torch.equal(collected1, hidden1) +def test_layer_hidden_collector_requires_target_layer_ids_from_config(monkeypatch): + _mock_target_layer_ids(monkeypatch, None) + model = SimpleNamespace(layers_num=2, pre_infer=_IdentityPreInfer()) + + with pytest.raises(AssertionError, match="target_layer_ids is required in draft config"): + LayerHiddenCollector(model=model) + + def test_noop_collector_keeps_normal_forward_output_minimal(): final_hidden = torch.randn(2, 3) collector = NoopHiddenCollector() collector.add(layer_index=0, hidden=final_hidden) + collector.add_final_hidden(final_hidden) - assert collector.prefill_outputs(final_hidden) == [final_hidden] - assert ( - collector.finish( - infer_state=None, - final_hidden=final_hidden, - ) - is None - ) + assert collector.finish(infer_state=None) is None def test_final_collector_returns_final_hidden_without_layer_bookkeeping(): final_hidden = torch.randn(2, 3) - collected = FinalHiddenCollector().finish( - infer_state=None, - final_hidden=final_hidden, - ) + collector = FinalHiddenCollector() + collector.add_final_hidden(final_hidden) + collected = collector.finish(infer_state=None) assert collected.data_ptr() == final_hidden.data_ptr() -def test_layer_collector_preserves_selected_layers_in_model_order(): +def test_layer_collector_preserves_selected_layers_in_model_order(monkeypatch): + _mock_target_layer_ids(monkeypatch, [0, 2]) layer0 = torch.full((2, 2), 1.0) layer1 = torch.full((2, 2), 2.0) layer2 = torch.full((2, 2), 3.0) model = SimpleNamespace(layers_num=3, pre_infer=_IdentityPreInfer()) - collector = LayerHiddenCollector(model=model, layer_ids=[0, 2]) + collector = LayerHiddenCollector(model=model) collector.add(layer_index=0, hidden=layer0) collector.add(layer_index=1, hidden=layer1) collector.add(layer_index=2, hidden=layer2) layer0.fill_(9.0) - forward_outputs = collector.prefill_outputs(layer2) - collected = collector.finish( - infer_state=SimpleNamespace(need_dp_prefill_balance=False), - final_hidden=layer2, - forward_outputs=forward_outputs, - ) + collected = collector.finish(infer_state=SimpleNamespace(need_dp_prefill_balance=False)) assert torch.equal(collected, torch.cat([torch.full((2, 2), 1.0), layer2], dim=-1)) assert not collector.layer_hiddens collector.add(layer_index=0, hidden=layer0) collector.add(layer_index=2, hidden=layer2) - collected = collector.finish( - infer_state=SimpleNamespace(need_dp_prefill_balance=False), - final_hidden=layer2, - ) + collected = collector.finish(infer_state=SimpleNamespace(need_dp_prefill_balance=False)) assert torch.equal(collected, torch.cat([layer0, layer2], dim=-1)) assert not collector.layer_hiddens + + +def test_layer_collector_restores_graph_state_without_sharing_runtime_container(monkeypatch): + _mock_target_layer_ids(monkeypatch, [0]) + model = SimpleNamespace(layers_num=2, pre_infer=_IdentityPreInfer()) + graph_collector = LayerHiddenCollector(model=model) + collector = graph_collector.new_instance() + infer_state = SimpleNamespace(need_dp_prefill_balance=False) + graph_collector.add(layer_index=0, hidden=torch.full((2, 3), 1.0)) + collector.restore_graph_state(graph_collector) + + collected = collector.finish(infer_state=infer_state) + + assert torch.equal(collected, torch.full((2, 3), 1.0)) + assert not collector.layer_hiddens + assert len(graph_collector.layer_hiddens) == 1 + + +def test_layer_collector_releases_graph_tensor_ownership_with_statistics(monkeypatch): + _mock_target_layer_ids(monkeypatch, [0, 2]) + model = SimpleNamespace(layers_num=3, pre_infer=_IdentityPreInfer()) + collector = LayerHiddenCollector(model=model) + hidden0 = torch.randn(2, 3) + hidden2 = torch.randn(2, 3) + converted = [] + + def to_no_ref(hidden): + converted.append(hidden) + return hidden + + monkeypatch.setattr(hidden_collector_module, "tensor_to_no_ref_tensor", to_no_ref) + collector.add(layer_index=0, hidden=hidden0) + collector.add(layer_index=2, hidden=hidden2) + + tensor_count, total_nbytes = collector.release_graph_tensor_ownership() + + assert tensor_count == 2 + assert total_nbytes == hidden0.numel() * hidden0.element_size() + hidden2.numel() * hidden2.element_size() + assert all(actual is expected for actual, expected in zip(converted, collector.layer_hiddens)) + assert NoopHiddenCollector().release_graph_tensor_ownership() == (0, 0) diff --git a/unit_tests/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py index 004d119f41..ce622000cd 100644 --- a/unit_tests/common/basemodel/test_mtp_manager.py +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -2,7 +2,9 @@ import pytest +import lightllm.common.basemodel.hidden_collector as hidden_collector_module import lightllm.common.basemodel.mtp_manager as mtp_manager_module +from lightllm.common.basemodel.hidden_collector import FinalHiddenCollector, LayerHiddenCollector, NoopHiddenCollector from lightllm.common.basemodel.mtp_manager import MtpManager @@ -84,3 +86,38 @@ def test_get_instance_returns_singleton(monkeypatch): monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) assert MtpManager.get_instance() is MtpManager.get_instance() + + +@pytest.mark.parametrize( + "spec_mode,is_draft_model,expected_type", + [ + (None, False, NoopHiddenCollector), + ("vanilla_with_att", False, FinalHiddenCollector), + ("eagle3", False, LayerHiddenCollector), + ("dspark", False, LayerHiddenCollector), + ("eagle3", True, FinalHiddenCollector), + ("dspark", True, NoopHiddenCollector), + ], +) +def test_create_hidden_collector_selects_implementation(monkeypatch, spec_mode, is_draft_model, expected_type): + args = SimpleNamespace( + mtp_mode=spec_mode, + mtp_step=7, + mtp_dynamic_verify=False, + mtp_draft_model_dir=["/models/draft"], + ) + monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) + monkeypatch.setattr(hidden_collector_module, "get_env_start_args", lambda: args) + monkeypatch.setattr( + hidden_collector_module.PretrainedConfig, + "get_config_dict", + lambda _: ({"target_layer_ids": [0]}, {}), + ) + model = SimpleNamespace(is_mtp_draft_model=is_draft_model, layers_num=2) + + prototype = MtpManager.get_instance().create_hidden_collector(model=model) + collector = prototype.new_instance() + + assert isinstance(prototype, expected_type) + assert isinstance(collector, expected_type) + assert collector is not prototype diff --git a/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py b/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py new file mode 100644 index 0000000000..3985e63de8 --- /dev/null +++ b/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py @@ -0,0 +1,83 @@ +from types import SimpleNamespace + +import torch + +import lightllm.common.basemodel.hidden_collector as hidden_collector_module +from lightllm.common.basemodel.hidden_collector import LayerHiddenCollector +from lightllm.common.basemodel.prefill_cuda_graph import PrefillCudaGraph + + +class _GraphInferState: + def __init__(self, hidden_collector): + self.hidden_collector = hidden_collector + self.copied_from = None + self.replayed_with = None + + def copy_for_prefill_cuda_graph(self, new_infer_state): + self.copied_from = new_infer_state + + def prefill_replay(self, new_infer_state): + self.replayed_with = new_infer_state + + +class _IdentityPreInfer: + @staticmethod + def _tpsp_allgather(input, infer_state): + del infer_state + return input + + +def _create_layer_collector(monkeypatch): + monkeypatch.setattr( + hidden_collector_module, + "get_env_start_args", + lambda: SimpleNamespace(mtp_draft_model_dir=["/models/draft"]), + ) + monkeypatch.setattr( + hidden_collector_module.PretrainedConfig, + "get_config_dict", + lambda _: ({"target_layer_ids": [0]}, {}), + ) + model = SimpleNamespace(layers_num=2, pre_infer=_IdentityPreInfer()) + return LayerHiddenCollector(model=model) + + +def test_replay_restores_captured_hidden_state_into_request_collector(monkeypatch): + graph_collector = _create_layer_collector(monkeypatch) + graph_collector.add(layer_index=0, hidden=torch.randn(2, 3)) + request_collector = graph_collector.new_instance() + graph_infer_state = _GraphInferState(hidden_collector=graph_collector) + request_infer_state = SimpleNamespace( + total_token_num=4, + prefix_total_token_num=0, + hidden_collector=request_collector, + ) + graph_output = torch.randn(2, 3) + prefill_graph = PrefillCudaGraph.__new__(PrefillCudaGraph) + prefill_graph.graph = {4: (graph_infer_state, [], [graph_output], graph_collector)} + + outputs = prefill_graph._replay(input_tensors=[], infer_state=request_infer_state) + + assert len(outputs) == 1 + assert outputs[0] is graph_output + assert request_infer_state.hidden_collector is request_collector + assert request_collector.layer_hiddens is not graph_collector.layer_hiddens + assert request_collector.layer_hiddens[0] is graph_collector.layer_hiddens[0] + assert graph_infer_state.copied_from is request_infer_state + assert graph_infer_state.replayed_with is request_infer_state + + +def test_first_replay_replaces_capture_collector_with_runtime_instance(monkeypatch): + graph_collector = _create_layer_collector(monkeypatch) + graph_collector.add(layer_index=0, hidden=torch.randn(2, 3)) + graph_infer_state = _GraphInferState(hidden_collector=graph_collector) + graph_infer_state.total_token_num = 4 + graph_infer_state.prefix_total_token_num = 0 + graph_output = torch.randn(2, 3) + prefill_graph = PrefillCudaGraph.__new__(PrefillCudaGraph) + prefill_graph.graph = {4: (graph_infer_state, [], [graph_output], graph_collector)} + + prefill_graph._replay(input_tensors=[], infer_state=graph_infer_state) + + assert graph_infer_state.hidden_collector is not graph_collector + assert graph_infer_state.hidden_collector.layer_hiddens[0] is graph_collector.layer_hiddens[0] diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 24612cf600..a12a2117a1 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -8,7 +8,6 @@ import lightllm.common.basemodel.attention.base_att as base_att_module import lightllm.common.basemodel.hidden_collector as hidden_collector_module from lightllm.common.basemodel.attention.base_att import BaseAttBackend -from lightllm.common.basemodel.hidden_collector import HiddenCollector from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer @@ -254,10 +253,16 @@ def test_dflash_expands_position_delta_with_request_block_rows(): def test_hidden_collector_reads_target_layer_ids(monkeypatch): + config_reads = [] + + def get_config_dict(path): + config_reads.append(path) + return {"target_layer_ids": [1, 20, 36]}, {} + monkeypatch.setattr( hidden_collector_module.PretrainedConfig, "get_config_dict", - lambda _: ({"target_layer_ids": [1, 20, 36]}, {}), + get_config_dict, ) monkeypatch.setattr( hidden_collector_module, @@ -266,9 +271,12 @@ def test_hidden_collector_reads_target_layer_ids(monkeypatch): ) model = SimpleNamespace(is_mtp_draft_model=False, layers_num=40) - collector = HiddenCollector(model=model, spec_mode="dspark") + collector = hidden_collector_module.LayerHiddenCollector(model=model) + new_collector = collector.new_instance() - assert collector.collectors[0].layer_ids == frozenset((1, 20, 36)) + assert collector.layer_ids == frozenset((1, 20, 36)) + assert new_collector.layer_ids == collector.layer_ids + assert config_reads == ["/models/dspark"] def test_dflash_added_kv_layers_come_from_draft_config(tmp_path): From 249667f24be56a4d1ad77bba6f6e7f699876dfcc Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 06:30:09 +0000 Subject: [PATCH 015/103] refactor: reuse forward in autotune warmup --- lightllm/common/basemodel/basemodel.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index d4a036502b..73aac4dff5 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -1204,9 +1204,7 @@ def _autotune_warmup(self): multimodal_params=[{"images": [], "audios": []}], **self._gen_special_model_input(total_token_num), ) - model_input.to_cuda() - assert model_input.mem_indexes.is_cuda - model_output = self._prefill(model_input=model_input) + model_output = self.forward(model_input) del model_output self.req_manager.free_all() self.mem_manager.free_all() From 55cf376ae361a35c4883acd7dfd089640c9e299d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 06:31:41 +0000 Subject: [PATCH 016/103] fix --- lightllm/common/basemodel/basemodel.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 73aac4dff5..e62d2e4440 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -1204,7 +1204,9 @@ def _autotune_warmup(self): multimodal_params=[{"images": [], "audios": []}], **self._gen_special_model_input(total_token_num), ) - model_output = self.forward(model_input) + model_output = self.forward( + model_input, + ) del model_output self.req_manager.free_all() self.mem_manager.free_all() From 56421dbbfbac8a51b7de2d3895297b74d88771d5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 06:50:57 +0000 Subject: [PATCH 017/103] refactor: initialize decode group metadata on demand --- lightllm/common/basemodel/batch_objs.py | 29 ++++++- lightllm/common/basemodel/cuda_graph.py | 4 - .../common/basemodel/test_model_input.py | 87 +++++++++++++++++++ 3 files changed, 113 insertions(+), 7 deletions(-) create mode 100644 unit_tests/common/basemodel/test_model_input.py diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index cd103ec216..db901f5a6d 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -1,7 +1,13 @@ +from dataclasses import dataclass +from typing import List, Optional + import torch -from dataclasses import dataclass, field -from typing import Optional -from typing import List + +from lightllm.utils.envs_utils import ( + enable_diverse_mode_gqa_decode_fast_kernel, + enable_triton_mtp_kernel, + get_env_start_args, +) from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor @@ -72,11 +78,13 @@ def to_cuda(self): self.b_ready_cache_len = self.b_ready_cache_len.cuda(non_blocking=True) if self.b_position_delta is not None: self.b_position_delta = self.b_position_delta.cuda(non_blocking=True) + assert self.is_prefill is False, "b_position_delta should only be used in decode phase." else: assert self.is_prefill is True, "decode ModelInput should provide b_position_delta." if self.b_prefill_start_loc is not None: self.b_prefill_start_loc = self.b_prefill_start_loc.cuda(non_blocking=True) + self._ensure_decode_group_metadata() if self.b_mark_shared_group is not None: self.b_mark_shared_group = self.b_mark_shared_group.cuda(non_blocking=True) if self.b_shared_seq_len is not None: @@ -92,6 +100,21 @@ def check_input(self): self.input_ids.dtype == torch.int64 ), f"model input_ids must use torch.int64, got {self.input_ids.dtype}" + def _ensure_decode_group_metadata(self): + """按 decode 模式补齐能够安全降级的 attention 分组信息。""" + if self.is_prefill: + return + + if enable_diverse_mode_gqa_decode_fast_kernel(): + # 缺少共享信息时退化为互不共享前缀的单行组,不改变 attention 结果。 + if self.b_mark_shared_group is None: + self.b_mark_shared_group = torch.ones_like(self.b_req_idx, dtype=torch.int32) + if self.b_shared_seq_len is None: + self.b_shared_seq_len = torch.zeros_like(self.b_req_idx, dtype=torch.int32) + elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): + if self.b_mark_shared_group is None: + self.b_mark_shared_group = torch.ones_like(self.b_req_idx, dtype=torch.int32) + @dataclass class ModelOutput: diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 58e2f90bb2..ad77f3aa70 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -258,7 +258,6 @@ def warmup(self, model): ) b_seq_len = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda") b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") - b_mark_shared_group = torch.zeros(batch_size, dtype=torch.int32, device="cuda") model_input = ModelInput( batch_size=batch_size, @@ -270,7 +269,6 @@ def warmup(self, model): b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, - b_mark_shared_group=b_mark_shared_group, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], @@ -317,7 +315,6 @@ def warmup_overlap(self, model): ) b_seq_len = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda") b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") - b_mark_shared_group = torch.zeros(batch_size, dtype=torch.int32, device="cuda") micro_batch = ModelInput( is_prefill=False, @@ -327,7 +324,6 @@ def warmup_overlap(self, model): max_kv_seq_len=max_len_in_batch, input_ids=input_ids, b_mtp_index=b_mtp_index, - b_mark_shared_group=b_mark_shared_group, mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_seq_len=b_seq_len, diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py new file mode 100644 index 0000000000..caab276449 --- /dev/null +++ b/unit_tests/common/basemodel/test_model_input.py @@ -0,0 +1,87 @@ +from types import SimpleNamespace + +import pytest +import torch + +import lightllm.common.basemodel.batch_objs as batch_objs_module +from lightllm.common.basemodel.batch_objs import ModelInput + + +def _create_model_input(*, is_prefill=False, b_mtp_index=None): + batch_size = 2 + return ModelInput( + batch_size=batch_size, + total_token_num=batch_size, + max_q_seq_len=1, + max_kv_seq_len=1, + b_req_idx=torch.arange(batch_size, dtype=torch.int32), + b_mtp_index=torch.zeros(batch_size, dtype=torch.int32) if b_mtp_index is None else b_mtp_index, + b_seq_len=torch.ones(batch_size, dtype=torch.int32), + is_prefill=is_prefill, + multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], + ) + + +def _mock_group_modes(monkeypatch, *, diverse=False, dynamic=False, triton=False): + monkeypatch.setattr( + batch_objs_module, + "enable_diverse_mode_gqa_decode_fast_kernel", + lambda: diverse, + ) + monkeypatch.setattr(batch_objs_module, "enable_triton_mtp_kernel", lambda: triton) + monkeypatch.setattr( + batch_objs_module, + "get_env_start_args", + lambda: SimpleNamespace(mtp_dynamic_verify=dynamic), + ) + + +def test_diverse_decode_defaults_to_independent_groups(monkeypatch): + _mock_group_modes(monkeypatch, diverse=True) + model_input = _create_model_input() + + model_input._ensure_decode_group_metadata() + + assert torch.equal(model_input.b_mark_shared_group, torch.ones(2, dtype=torch.int32)) + assert torch.equal(model_input.b_shared_seq_len, torch.zeros(2, dtype=torch.int32)) + + +@pytest.mark.parametrize("dynamic,triton", [(True, False), (False, True)]) +def test_plain_spec_decode_defaults_to_single_row_groups(monkeypatch, dynamic, triton): + _mock_group_modes(monkeypatch, dynamic=dynamic, triton=triton) + model_input = _create_model_input() + + model_input._ensure_decode_group_metadata() + + assert torch.equal(model_input.b_mark_shared_group, torch.ones(2, dtype=torch.int32)) + assert model_input.b_shared_seq_len is None + + +def test_multi_row_spec_decode_defaults_to_single_row_groups(monkeypatch): + _mock_group_modes(monkeypatch, dynamic=True) + model_input = _create_model_input(b_mtp_index=torch.tensor([0, 1], dtype=torch.int32)) + + model_input._ensure_decode_group_metadata() + + assert torch.equal(model_input.b_mark_shared_group, torch.ones(2, dtype=torch.int32)) + + +def test_multi_row_spec_decode_preserves_explicit_group_metadata(monkeypatch): + _mock_group_modes(monkeypatch, dynamic=True) + model_input = _create_model_input(b_mtp_index=torch.tensor([0, 1], dtype=torch.int32)) + group_markers = torch.tensor([0, 2], dtype=torch.int32) + model_input.b_mark_shared_group = group_markers + + model_input._ensure_decode_group_metadata() + + assert model_input.b_mark_shared_group is group_markers + + +def test_prefill_does_not_create_decode_group_metadata(monkeypatch): + _mock_group_modes(monkeypatch, diverse=True, dynamic=True, triton=True) + model_input = _create_model_input(is_prefill=True) + + model_input._ensure_decode_group_metadata() + + assert model_input.b_mark_shared_group is None + assert model_input.b_shared_seq_len is None From b25e38f81da7f7607a8e62b2d49d4d815f0d0c8e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 07:49:22 +0000 Subject: [PATCH 018/103] refactor: split FA3 decode state initialization --- .../common/basemodel/attention/base_att.py | 10 +-- lightllm/common/basemodel/attention/fa3/fp.py | 88 +++++++++++-------- .../common/basemodel/attention/fa3/mla.py | 88 +++++++++++-------- unit_tests/utils/test_speculative_utils.py | 2 + 4 files changed, 106 insertions(+), 82 deletions(-) diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index 00176b973f..6d1bf55f50 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -42,12 +42,10 @@ def create_att_decode_state(self) -> "BaseDecodeAttState": def uses_dynamic_spec_verify_layout(self) -> bool: args = get_env_start_args() draft_step = self.model.mtp_manager.get_decode_draft_step(self.model.is_mtp_draft_model) - if draft_step == 0 or not args.mtp_dynamic_verify: - return False - - # Target verification may compact each request to a different row count. - # Parallel block drafter forwards still use their checkpoint-defined fixed layout. - return args.mtp_mode not in ("dspark", "dflash") or not self.model.is_mtp_draft_model + is_main_model = not self.model.is_mtp_draft_model + has_decode_draft_step = draft_step > 0 + dynamic_verify_enabled = args.mtp_dynamic_verify + return is_main_model and has_decode_draft_step and dynamic_verify_enabled def _find_layer_index( self, k: torch.Tensor, v: torch.Tensor, att_state: Union["BasePrefillAttState", "BaseDecodeAttState"] diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index 311e751732..c617db3c50 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -134,46 +134,60 @@ class Fa3DecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: Fa3AttBackend = self.backend draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) - decode_rows_per_request = draft_step + 1 - uses_dynamic_spec_verify_layout = self.backend.uses_dynamic_spec_verify_layout() - - if draft_step > 0 and not uses_dynamic_spec_verify_layout: - assert self.infer_state.batch_size % decode_rows_per_request == 0, ( - "FA3 fixed-layout decode requires batch_size to be divisible by draft_step + 1, " - f"got batch_size={self.infer_state.batch_size}, draft_step={draft_step}." - ) - - # 修正 mtp 在 fa3 下的输入。 - if uses_dynamic_spec_verify_layout: - (b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len,) = build_dynamic_spec_fa3_decode_params( - b_req_idx=self.infer_state.b_req_idx, - b_seq_len=self.infer_state.b_seq_len, - b_mark_shared_group=self.infer_state.b_mark_shared_group, - att_batch_size=self.infer_state.batch_size, - hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, - ) + if self.backend.uses_dynamic_spec_verify_layout(): + b_att_req_idx = self._init_dynamic_spec_verify_state(draft_step) elif draft_step > 0: - b_q_seq_len = torch.full( - (self.infer_state.b_seq_len.shape[0] // decode_rows_per_request,), - fill_value=decode_rows_per_request, - dtype=torch.int32, - device=self.infer_state.b_seq_len.device, - ) - b_kv_seq_len = self.infer_state.b_seq_len[draft_step::decode_rows_per_request] - b_att_req_idx = self.infer_state.b_req_idx[draft_step::decode_rows_per_request] - self.b_att_seq_len = b_kv_seq_len.contiguous() + b_att_req_idx = self._init_fixed_spec_decode_state(draft_step) else: - b_att_req_idx = self.infer_state.b_req_idx - self.b_att_seq_len = self.infer_state.b_seq_len + b_att_req_idx = self._init_normal_decode_state() - if draft_step > 0: - b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) - self.cu_seqlens_q = b1_cu_q_seq_len.int() - self.cu_seqlens_k = b1_cu_kv_seq_len.int() - else: - self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() - self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() + self._init_page_table(b_att_req_idx) + + def _init_dynamic_spec_verify_state(self, draft_step: int) -> torch.Tensor: + b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len = build_dynamic_spec_fa3_decode_params( + b_req_idx=self.infer_state.b_req_idx, + b_seq_len=self.infer_state.b_seq_len, + b_mark_shared_group=self.infer_state.b_mark_shared_group, + att_batch_size=self.infer_state.batch_size, + hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, + ) + self._init_spec_decode_cu_seqlens(b_q_seq_len, b_kv_seq_len) + self.decode_max_q_seq_len = draft_step + 1 + return b_att_req_idx + + def _init_fixed_spec_decode_state(self, draft_step: int) -> torch.Tensor: + mtp_size = draft_step + 1 + assert self.infer_state.batch_size % mtp_size == 0, ( + "FA3 fixed-layout decode requires batch_size to be divisible by draft_step + 1, " + f"got batch_size={self.infer_state.batch_size}, draft_step={draft_step}." + ) + + b_q_seq_len = torch.full( + (self.infer_state.b_seq_len.shape[0] // mtp_size,), + fill_value=mtp_size, + dtype=torch.int32, + device=self.infer_state.b_seq_len.device, + ) + b_kv_seq_len = self.infer_state.b_seq_len[draft_step::mtp_size] + b_att_req_idx = self.infer_state.b_req_idx[draft_step::mtp_size] + self.b_att_seq_len = b_kv_seq_len.contiguous() + self._init_spec_decode_cu_seqlens(b_q_seq_len, b_kv_seq_len) + self.decode_max_q_seq_len = mtp_size + return b_att_req_idx + + def _init_normal_decode_state(self) -> torch.Tensor: + self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() + self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() + self.b_att_seq_len = self.infer_state.b_seq_len + self.decode_max_q_seq_len = 1 + return self.infer_state.b_req_idx + + def _init_spec_decode_cu_seqlens(self, b_q_seq_len: torch.Tensor, b_kv_seq_len: torch.Tensor): + b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) + self.cu_seqlens_q = b1_cu_q_seq_len.int() + self.cu_seqlens_k = b1_cu_kv_seq_len.int() + def _init_page_table(self, b_att_req_idx: torch.Tensor): att_batch_size = b_att_req_idx.shape[0] model = self.backend.model # 可以使用 cuda graph的时候从 buffer中申请 @@ -197,8 +211,6 @@ def init_state(self): req_to_token_indexs=model.req_manager.req_to_token_indexs, b_req_idx=b_att_req_idx, ) - self.decode_max_q_seq_len = decode_rows_per_request - return def copy_for_decode_cuda_graph(self, new_state: "Fa3DecodeAttState"): super().copy_for_decode_cuda_graph(new_state) diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index bf32efb7b4..0cd82accf2 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -114,46 +114,60 @@ class MlaFa3DecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: MlaFa3AttBackend = self.backend draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) - decode_rows_per_request = draft_step + 1 - uses_dynamic_spec_verify_layout = self.backend.uses_dynamic_spec_verify_layout() - - if draft_step > 0 and not uses_dynamic_spec_verify_layout: - assert self.infer_state.batch_size % decode_rows_per_request == 0, ( - "FA3 fixed-layout decode requires batch_size to be divisible by draft_step + 1, " - f"got batch_size={self.infer_state.batch_size}, draft_step={draft_step}." - ) - - # 修正 mtp 在 fa3 下的输入。 - if uses_dynamic_spec_verify_layout: - (b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len,) = build_dynamic_spec_fa3_decode_params( - b_req_idx=self.infer_state.b_req_idx, - b_seq_len=self.infer_state.b_seq_len, - b_mark_shared_group=self.infer_state.b_mark_shared_group, - att_batch_size=self.infer_state.batch_size, - hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, - ) + if self.backend.uses_dynamic_spec_verify_layout(): + b_att_req_idx = self._init_dynamic_spec_verify_state(draft_step) elif draft_step > 0: - b_q_seq_len = torch.full( - (self.infer_state.b_seq_len.shape[0] // decode_rows_per_request,), - fill_value=decode_rows_per_request, - dtype=torch.int32, - device=self.infer_state.b_seq_len.device, - ) - b_kv_seq_len = self.infer_state.b_seq_len[draft_step::decode_rows_per_request] - b_att_req_idx = self.infer_state.b_req_idx[draft_step::decode_rows_per_request] - self.b_att_seq_len = b_kv_seq_len.contiguous() + b_att_req_idx = self._init_fixed_spec_decode_state(draft_step) else: - b_att_req_idx = self.infer_state.b_req_idx - self.b_att_seq_len = self.infer_state.b_seq_len + b_att_req_idx = self._init_normal_decode_state() - if draft_step > 0: - b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) - self.cu_seqlens_q = b1_cu_q_seq_len.int() - self.cu_seqlens_k = b1_cu_kv_seq_len.int() - else: - self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() - self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() + self._init_page_table(b_att_req_idx) + + def _init_dynamic_spec_verify_state(self, draft_step: int) -> torch.Tensor: + b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len = build_dynamic_spec_fa3_decode_params( + b_req_idx=self.infer_state.b_req_idx, + b_seq_len=self.infer_state.b_seq_len, + b_mark_shared_group=self.infer_state.b_mark_shared_group, + att_batch_size=self.infer_state.batch_size, + hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, + ) + self._init_spec_decode_cu_seqlens(b_q_seq_len, b_kv_seq_len) + self.decode_max_q_seq_len = draft_step + 1 + return b_att_req_idx + + def _init_fixed_spec_decode_state(self, draft_step: int) -> torch.Tensor: + mtp_size = draft_step + 1 + assert self.infer_state.batch_size % mtp_size == 0, ( + "FA3 fixed-layout decode requires batch_size to be divisible by draft_step + 1, " + f"got batch_size={self.infer_state.batch_size}, draft_step={draft_step}." + ) + + b_q_seq_len = torch.full( + (self.infer_state.b_seq_len.shape[0] // mtp_size,), + fill_value=mtp_size, + dtype=torch.int32, + device=self.infer_state.b_seq_len.device, + ) + b_kv_seq_len = self.infer_state.b_seq_len[draft_step::mtp_size] + b_att_req_idx = self.infer_state.b_req_idx[draft_step::mtp_size] + self.b_att_seq_len = b_kv_seq_len.contiguous() + self._init_spec_decode_cu_seqlens(b_q_seq_len, b_kv_seq_len) + self.decode_max_q_seq_len = mtp_size + return b_att_req_idx + + def _init_normal_decode_state(self) -> torch.Tensor: + self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() + self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() + self.b_att_seq_len = self.infer_state.b_seq_len + self.decode_max_q_seq_len = 1 + return self.infer_state.b_req_idx + + def _init_spec_decode_cu_seqlens(self, b_q_seq_len: torch.Tensor, b_kv_seq_len: torch.Tensor): + b1_cu_q_seq_len, b1_cu_kv_seq_len = gen_cumsum_pad0_tensor(b_q_seq_len, b_kv_seq_len) + self.cu_seqlens_q = b1_cu_q_seq_len.int() + self.cu_seqlens_k = b1_cu_kv_seq_len.int() + def _init_page_table(self, b_att_req_idx: torch.Tensor): att_batch_size = b_att_req_idx.shape[0] model = self.backend.model # 可以使用 cuda graph的时候从 buffer中申请 @@ -177,8 +191,6 @@ def init_state(self): req_to_token_indexs=model.req_manager.req_to_token_indexs, b_req_idx=b_att_req_idx, ) - self.decode_max_q_seq_len = decode_rows_per_request - return def copy_for_decode_cuda_graph(self, new_state: "MlaFa3DecodeAttState"): super().copy_for_decode_cuda_graph(new_state) diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index a12a2117a1..b3f2daafd3 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -52,8 +52,10 @@ def test_qwen3_eagle_uses_layers_checkpoint_prefix(): ("dspark", True, 7, True, False), ("dflash", False, 7, True, True), ("dflash", True, 7, True, False), + ("vanilla_with_att", True, 7, True, False), ("vanilla_with_att", True, 0, True, False), ("eagle3", True, 0, True, False), + ("eagle3", False, 7, True, True), ("eagle3", False, 7, False, False), ], ) From 8f32f616790081c28fb17e7667877fce7d6a7ed5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 08:04:29 +0000 Subject: [PATCH 019/103] refactor: move FA3 causality into attention state --- .../common/basemodel/attention/base_att.py | 5 ++ lightllm/common/basemodel/attention/fa3/fp.py | 8 ++- .../common/basemodel/attention/fa3/fp8.py | 4 +- .../common/basemodel/attention/fa3/mla.py | 8 ++- lightllm/common/basemodel/infer_struct.py | 3 - lightllm/models/qwen3_dflash/infer_struct.py | 7 -- lightllm/models/qwen3_dflash/model.py | 2 +- unit_tests/utils/test_speculative_utils.py | 72 +++++++++++++++++++ 8 files changed, 92 insertions(+), 17 deletions(-) diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index 6d1bf55f50..6dd8d2c368 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -47,6 +47,11 @@ def uses_dynamic_spec_verify_layout(self) -> bool: dynamic_verify_enabled = args.mtp_dynamic_verify return is_main_model and has_decode_draft_step and dynamic_verify_enabled + def uses_causal_attention(self) -> bool: + args = get_env_start_args() + is_parallel_block_draft = self.model.is_mtp_draft_model and args.mtp_mode in ("dspark", "dflash") + return not is_parallel_block_draft + def _find_layer_index( self, k: torch.Tensor, v: torch.Tensor, att_state: Union["BasePrefillAttState", "BaseDecodeAttState"] ) -> int: diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index c617db3c50..be190a05bc 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -51,8 +51,10 @@ class Fa3PrefillAttState(BasePrefillAttState): cu_seqlens_q: torch.Tensor = None cu_seqlens_k: torch.Tensor = None page_table: torch.Tensor = None + causal: bool = None def init_state(self): + self.causal = self.backend.uses_causal_attention() self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() self.page_table = torch.empty( @@ -111,7 +113,7 @@ def _nomarl_prefill_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.infer_state.max_q_seq_len, softmax_scale=sm_scale, - causal=self.infer_state.prefill_causal, + causal=self.causal, window_size=window_size, softcap=0.0, k_descale=k_descale, @@ -130,9 +132,11 @@ class Fa3DecodeAttState(BaseDecodeAttState): b_att_seq_len: torch.Tensor = None # 在是否开启mtp 的不同模式下,其设置不同的值,可以加速算子的运行。 decode_max_q_seq_len: int = None + causal: bool = None def init_state(self): self.backend: Fa3AttBackend = self.backend + self.causal = self.backend.uses_causal_attention() draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) if self.backend.uses_dynamic_spec_verify_layout(): b_att_req_idx = self._init_dynamic_spec_verify_state(draft_step) @@ -263,7 +267,7 @@ def _normal_decode_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, softmax_scale=sm_scale, - causal=self.infer_state.decode_causal, + causal=self.causal, window_size=window_size, softcap=0.0, k_descale=k_descale, diff --git a/lightllm/common/basemodel/attention/fa3/fp8.py b/lightllm/common/basemodel/attention/fa3/fp8.py index dea9ed1e0d..6e6d8a8366 100644 --- a/lightllm/common/basemodel/attention/fa3/fp8.py +++ b/lightllm/common/basemodel/attention/fa3/fp8.py @@ -98,7 +98,7 @@ def _fp8_prefill_att( cu_seqlens_q=self.cu_seqlens_q, cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.infer_state.max_q_seq_len, - causal=self.infer_state.prefill_causal, + causal=self.causal, window_size=(-1, -1), softcap=0.0, q_descale=q_scale, @@ -185,7 +185,7 @@ def _fp8_decode_att( cu_seqlens_q=self.cu_seqlens_q, cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, - causal=self.infer_state.decode_causal, + causal=self.causal, window_size=(-1, -1), softcap=0.0, q_descale=q_scale.view(self.infer_state.batch_size, k_head_num), diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index 0cd82accf2..8805c6cae4 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -48,8 +48,10 @@ def create_att_decode_state(self, infer_state) -> "MlaFa3DecodeAttState": class MlaFa3PrefillAttState(BasePrefillAttState): cu_seqlens_q: torch.Tensor = None cu_seqlens_k: torch.Tensor = None + causal: bool = None def init_state(self): + self.causal = self.backend.uses_causal_attention() self.cu_seqlens_q = self.infer_state.b1_cu_q_seq_len.int() self.cu_seqlens_k = self.infer_state.b1_cu_kv_seq_len.int() @@ -96,7 +98,7 @@ def _mla_prefill_att( max_seqlen_q=self.infer_state.max_q_seq_len, max_seqlen_k=self.infer_state.max_kv_seq_len, softmax_scale=softmax_scale, - causal=self.infer_state.prefill_causal, + causal=self.causal, return_softmax_lse=False, ) return o_tensor @@ -110,9 +112,11 @@ class MlaFa3DecodeAttState(BaseDecodeAttState): b_att_seq_len: torch.Tensor = None # 在是否开启mtp 的不同模式下,其设置不同的值,可以加速算子的运行。 decode_max_q_seq_len: int = None + causal: bool = None def init_state(self): self.backend: MlaFa3AttBackend = self.backend + self.causal = self.backend.uses_causal_attention() draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) if self.backend.uses_dynamic_spec_verify_layout(): b_att_req_idx = self._init_dynamic_spec_verify_state(draft_step) @@ -247,7 +251,7 @@ def _mla_decode_att( cu_seqlens_k_new=self.cu_seqlens_k, max_seqlen_q=self.decode_max_q_seq_len, softmax_scale=softmax_scale, - causal=self.infer_state.decode_causal, + causal=self.causal, window_size=(-1, -1), softcap=0.0, k_descale=k_descale, diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index 5fb5599469..d4835d2fc6 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -48,9 +48,6 @@ def __init__(self): # 的sum值, 其值等于 sum(b_ready_cache_len) self.prefix_total_token_num: int = None self.is_prefill: bool = None - # Specialized models override attention causality explicitly. - self.prefill_causal: bool = True - self.decode_causal: bool = True self.mem_manager: MemoryManager = None self.req_manager: ReqManager = None diff --git a/lightllm/models/qwen3_dflash/infer_struct.py b/lightllm/models/qwen3_dflash/infer_struct.py index a391089087..b33aac3179 100644 --- a/lightllm/models/qwen3_dflash/infer_struct.py +++ b/lightllm/models/qwen3_dflash/infer_struct.py @@ -3,10 +3,3 @@ class Qwen3DFlashInferStateInfo(LlamaInferStateInfo): """DFlash attention metadata.""" - - def init_some_extra_state(self, model): - super().init_some_extra_state(model) - if self.is_prefill: - self.prefill_causal = False - else: - self.decode_causal = False diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py index 1f7ebe644d..e0df32bfd3 100644 --- a/lightllm/models/qwen3_dflash/model.py +++ b/lightllm/models/qwen3_dflash/model.py @@ -53,7 +53,7 @@ def _init_mem_manager(self): def _init_att_backend(self): super()._init_att_backend() - # FA3 is currently the only backend that honors decode_causal=False. + # FA3 is currently the only backend that supports non-causal block decode. # TODO: Remove this restriction after Triton and FlashInfer support non-causal block decode. if not isinstance(self.decode_att_backend, (Fa3AttBackend, Fp8Fa3AttBackend)): raise NotImplementedError("Qwen3 DFlash decode requires FA3") diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index b3f2daafd3..f2fa663e7e 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -8,6 +8,8 @@ import lightllm.common.basemodel.attention.base_att as base_att_module import lightllm.common.basemodel.hidden_collector as hidden_collector_module from lightllm.common.basemodel.attention.base_att import BaseAttBackend +from lightllm.common.basemodel.attention.fa3.fp import Fa3DecodeAttState, Fa3PrefillAttState +from lightllm.common.basemodel.attention.fa3.mla import MlaFa3DecodeAttState, MlaFa3PrefillAttState from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer @@ -82,6 +84,76 @@ def test_attention_backend_selects_dynamic_spec_layout( assert BaseAttBackend.uses_dynamic_spec_verify_layout(backend) is expected +@pytest.mark.parametrize( + "mtp_mode, is_draft_model, expected", + [ + (None, False, True), + ("dflash", False, True), + ("dflash", True, False), + ("dspark", True, False), + ("eagle3", True, True), + ("vanilla_with_att", True, True), + ], +) +def test_attention_backend_selects_causality(monkeypatch, mtp_mode, is_draft_model, expected): + monkeypatch.setattr( + base_att_module, + "get_env_start_args", + lambda: SimpleNamespace(mtp_mode=mtp_mode), + ) + backend = SimpleNamespace(model=SimpleNamespace(is_mtp_draft_model=is_draft_model)) + + assert BaseAttBackend.uses_causal_attention(backend) is expected + + +@pytest.mark.parametrize("state_class", [Fa3PrefillAttState, MlaFa3PrefillAttState]) +def test_fa3_prefill_state_owns_causality(state_class): + infer_state = SimpleNamespace( + b1_cu_q_seq_len=torch.tensor([0, 1, 2], dtype=torch.int32), + b1_cu_kv_seq_len=torch.tensor([0, 3, 7], dtype=torch.int32), + b_req_idx=torch.tensor([0, 1], dtype=torch.int32), + batch_size=2, + max_kv_seq_len=4, + input_ids=torch.empty(2, dtype=torch.int64), + req_manager=SimpleNamespace(req_to_token_indexs=torch.arange(8, dtype=torch.int32).reshape(2, 4)), + ) + state = state_class( + backend=SimpleNamespace(uses_causal_attention=lambda: False), + infer_state=infer_state, + ) + + state.init_state() + + assert state.causal is False + + +@pytest.mark.parametrize("state_class", [Fa3DecodeAttState, MlaFa3DecodeAttState]) +def test_fa3_decode_state_owns_causality(state_class): + infer_state = SimpleNamespace( + b1_cu_q_seq_len=torch.tensor([0, 1, 2], dtype=torch.int32), + b1_cu_kv_seq_len=torch.tensor([0, 3, 7], dtype=torch.int32), + b_req_idx=torch.tensor([0, 1], dtype=torch.int32), + b_seq_len=torch.tensor([3, 4], dtype=torch.int32), + ) + model = SimpleNamespace( + is_mtp_draft_model=False, + mtp_manager=SimpleNamespace(get_decode_draft_step=lambda _: 0), + ) + state = state_class( + backend=SimpleNamespace( + model=model, + uses_causal_attention=lambda: False, + uses_dynamic_spec_verify_layout=lambda: False, + ), + infer_state=infer_state, + ) + state._init_page_table = lambda _: None + + state.init_state() + + assert state.causal is False + + @pytest.mark.parametrize( "model_type, spec_mode, expected_class_name", [ From 538189a8367d212be3548a17ca7cf0fc6a83842b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 08:24:34 +0000 Subject: [PATCH 020/103] refactor: align linear attention MTP state naming --- .../common/basemodel/attention/linear/gdn.py | 95 +++++++++++-------- .../triton_kernel/linear_att/__init__.py | 2 +- ...al_conv1d_spec.py => causal_conv1d_mtp.py} | 8 +- .../linear_att/mtp_fused_recurrent.py | 2 +- ...ec_state_params.py => mtp_state_params.py} | 8 +- .../basemodel/attention/linear/test_gdn.py | 95 +++++++++++++++++++ ...nv1d_spec.py => test_causal_conv1d_mtp.py} | 14 +-- 7 files changed, 168 insertions(+), 56 deletions(-) rename lightllm/common/basemodel/triton_kernel/linear_att/{causal_conv1d_spec.py => causal_conv1d_mtp.py} (98%) rename lightllm/common/basemodel/triton_kernel/linear_att/{spec_state_params.py => mtp_state_params.py} (94%) rename unit_tests/common/basemodel/triton_kernel/linear_att/{test_causal_conv1d_spec.py => test_causal_conv1d_mtp.py} (98%) diff --git a/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index 44cb08815e..872920d779 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -10,8 +10,8 @@ from lightllm.common.basemodel.triton_kernel.linear_att.mtp_fused_recurrent import ( mtp_fused_recurrent_gated_delta_rule, ) -from lightllm.common.basemodel.triton_kernel.linear_att.spec_state_params import ( - build_dynamic_spec_linear_att_state_params, +from lightllm.common.basemodel.triton_kernel.linear_att.mtp_state_params import ( + build_dynamic_mtp_linear_att_state_params, ) from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import ( fused_recurrent_gated_delta_rule, @@ -202,46 +202,63 @@ class LinearAttDecodeAttState(BaseDecodeAttState): b_conv_buffer_idx: torch.Tensor = None b_ssm_buffer_idx: torch.Tensor = None - b1_spec_cu_q_seq_len: torch.Tensor = None + b1_mtp_cu_q_seq_len: torch.Tensor = None b_num_accepted_tokens: torch.Tensor = None def init_state(self): draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) - if draft_step == 0: - self.b_conv_buffer_idx = self.infer_state.b_req_idx - self.b_ssm_buffer_idx = self.infer_state.b_req_idx - return + self._init_normal_decode_state() + elif self.backend.uses_dynamic_spec_verify_layout(): + self._init_dynamic_mtp_decode_state(draft_step + 1) + else: + self._init_fixed_mtp_decode_state(draft_step) + + def _init_normal_decode_state(self): + self.b_conv_buffer_idx = self.infer_state.b_req_idx + self.b_ssm_buffer_idx = self.infer_state.b_req_idx + def _init_dynamic_mtp_decode_state(self, mtp_size: int): + ( + self.b1_mtp_cu_q_seq_len, + self.b_conv_buffer_idx, + self.b_num_accepted_tokens, + ) = build_dynamic_mtp_linear_att_state_params( + b_req_idx=self.infer_state.b_req_idx, + b_mtp_index=self.infer_state.b_mtp_index, + req_to_mtp_state_index=self.infer_state.req_manager.req_to_mtp_state_index, + hold_req_id=self.infer_state.req_manager.HOLD_REQUEST_ID, + ) + self._init_mtp_ssm_buffer_idx(mtp_size) + + def _init_fixed_mtp_decode_state(self, draft_step: int): + mtp_size = draft_step + 1 batch_size = self.infer_state.batch_size - device = self.infer_state.b_req_idx.device - if self.backend.uses_dynamic_spec_verify_layout(): - ( - self.b1_spec_cu_q_seq_len, - self.b_conv_buffer_idx, - self.b_num_accepted_tokens, - ) = build_dynamic_spec_linear_att_state_params( - b_req_idx=self.infer_state.b_req_idx, - b_mtp_index=self.infer_state.b_mtp_index, - req_to_mtp_state_index=self.infer_state.req_manager.req_to_mtp_state_index, - hold_req_id=self.infer_state.req_manager.HOLD_REQUEST_ID, - ) - else: - assert batch_size % (draft_step + 1) == 0, ( - "GDN fixed-layout decode requires batch_size to be divisible by draft_step + 1, " - f"got batch_size={batch_size}, draft_step={draft_step}." - ) - att_batch_size = batch_size // (draft_step + 1) - self.b1_spec_cu_q_seq_len = torch.arange( - 0, batch_size + 1, draft_step + 1, dtype=torch.int32, device=device - ) - self.b_conv_buffer_idx = self.infer_state.b_req_idx.view(att_batch_size, draft_step + 1)[:, 0].contiguous() - self.b_num_accepted_tokens = self.infer_state.req_manager.req_to_mtp_state_index[self.b_conv_buffer_idx] + 1 + assert batch_size % mtp_size == 0, ( + "GDN fixed-layout decode requires batch_size to be divisible by draft_step + 1, " + f"got batch_size={batch_size}, draft_step={draft_step}." + ) + att_batch_size = batch_size // mtp_size + self.b1_mtp_cu_q_seq_len = torch.arange( + 0, + batch_size + 1, + mtp_size, + dtype=torch.int32, + device=self.infer_state.b_req_idx.device, + ) + self.b_conv_buffer_idx = self.infer_state.b_req_idx.view(att_batch_size, mtp_size)[:, 0].contiguous() + self.b_num_accepted_tokens = self.infer_state.req_manager.req_to_mtp_state_index[self.b_conv_buffer_idx] + 1 + self._init_mtp_ssm_buffer_idx(mtp_size) + + def _init_mtp_ssm_buffer_idx(self, mtp_size: int): # Each request owns one recurrent-state slot per verify row. - state_offsets = torch.arange(draft_step + 1, device=device, dtype=self.infer_state.b_req_idx.dtype) - self.b_ssm_buffer_idx = self.b_conv_buffer_idx[:, None] * (draft_step + 1) + state_offsets[None, :] - return + state_offsets = torch.arange( + mtp_size, + device=self.infer_state.b_req_idx.device, + dtype=self.infer_state.b_req_idx.dtype, + ) + self.b_ssm_buffer_idx = self.b_conv_buffer_idx[:, None] * mtp_size + state_offsets[None, :] def decode_att( self, @@ -265,7 +282,7 @@ def decode_att( draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) if draft_step > 0: - core_attn_out = self._gdn_spec_kernel( + core_attn_out = self._gdn_mtp_kernel( mixed_qkv, conv_states, ssm_states, @@ -335,7 +352,7 @@ def _gdn_decode_kernel( ) return core_attn_out, z - def _gdn_spec_kernel( + def _gdn_mtp_kernel( self, mixed_qkv: torch.Tensor, conv_states: torch.Tensor, @@ -345,14 +362,14 @@ def _gdn_spec_kernel( infer_state: "Qwen3NextInferStateInfo", layer_weight: "Qwen3NextTransformerLayerWeight", ): - from lightllm.common.basemodel.triton_kernel.linear_att.causal_conv1d_spec import ( - causal_conv1d_update as causal_conv1d_update_spec, + from lightllm.common.basemodel.triton_kernel.linear_att.causal_conv1d_mtp import ( + causal_conv1d_update as causal_conv1d_update_mtp, ) backend: LinearAttBackend = self.backend - cu_seqlens_q = self.b1_spec_cu_q_seq_len - mixed_qkv = causal_conv1d_update_spec( + cu_seqlens_q = self.b1_mtp_cu_q_seq_len + mixed_qkv = causal_conv1d_update_mtp( mixed_qkv, conv_states, layer_weight.linear_conv1d.mm_param.weight, diff --git a/lightllm/common/basemodel/triton_kernel/linear_att/__init__.py b/lightllm/common/basemodel/triton_kernel/linear_att/__init__.py index 3114101b2b..aa7cfa9f4a 100644 --- a/lightllm/common/basemodel/triton_kernel/linear_att/__init__.py +++ b/lightllm/common/basemodel/triton_kernel/linear_att/__init__.py @@ -1,7 +1,7 @@ """Linear-attention / GDN triton kernels shared across hybrid models.""" from .causal_conv1d import causal_conv1d_fn -from .causal_conv1d_spec import causal_conv1d_update +from .causal_conv1d_mtp import causal_conv1d_update from .fused_gdn_gating import fused_gdn_gating from .gdn_decode_pack import conv_pack_gdn_decode_inputs from .mtp_fused_recurrent import mtp_fused_recurrent_gated_delta_rule diff --git a/lightllm/common/basemodel/triton_kernel/linear_att/causal_conv1d_spec.py b/lightllm/common/basemodel/triton_kernel/linear_att/causal_conv1d_mtp.py similarity index 98% rename from lightllm/common/basemodel/triton_kernel/linear_att/causal_conv1d_spec.py rename to lightllm/common/basemodel/triton_kernel/linear_att/causal_conv1d_mtp.py index 825a164447..5e5dde09ee 100644 --- a/lightllm/common/basemodel/triton_kernel/linear_att/causal_conv1d_spec.py +++ b/lightllm/common/basemodel/triton_kernel/linear_att/causal_conv1d_mtp.py @@ -5,8 +5,8 @@ # - imports point at standard triton instead of vLLM's triton-lite. # - vLLM block-table params (block_idx_last_scheduled_token, initial_state_idx, # null_block_id) are dropped; LightLLM uses contiguous per-request slots. -# - IS_VARLEN / IS_SPEC_DECODING / non-spec paths removed; this kernel now -# exclusively serves the spec-decode varlen path (with num_accepted_tokens, +# - Upstream non-MTP paths were removed; this kernel now exclusively serves +# the MTP decode varlen path (with num_accepted_tokens, # query_start_loc and mtp_step all required). # - One widened conv_state slot per request holds K-1+mtp_step positions. # The read offset is num_accepted_tokens-1; writes go back to the same slot. @@ -314,7 +314,7 @@ def causal_conv1d_update( query_start_loc: Optional[torch.Tensor] = None, pad_slot_id: int = -1, ): - """Spec-decode causal depthwise conv1d update. + """MTP decode causal depthwise conv1d update. Processes ``mtp_step + 1`` tokens per request in varlen layout. Uses a single widened conv_state slot per request that holds @@ -329,7 +329,7 @@ def causal_conv1d_update( conv_state: ``(num_slots, dim, state_len)`` float with ``state_len == width - 1 + mtp_step``. weight: depthwise filter of shape ``(dim, width)``. - mtp_step: number of speculative (draft) tokens per request + mtp_step: number of extra MTP tokens per request (``seqlen == mtp_step + 1``). bias: optional ``(dim,)`` float bias. activation: ``None``, ``"silu"`` or ``"swish"``. diff --git a/lightllm/common/basemodel/triton_kernel/linear_att/mtp_fused_recurrent.py b/lightllm/common/basemodel/triton_kernel/linear_att/mtp_fused_recurrent.py index e868387dc8..2eb5ef5333 100644 --- a/lightllm/common/basemodel/triton_kernel/linear_att/mtp_fused_recurrent.py +++ b/lightllm/common/basemodel/triton_kernel/linear_att/mtp_fused_recurrent.py @@ -4,7 +4,7 @@ # SPDX-FileCopyrightText: Songlin Yang, Yu Zhang # # Extracted from fused_recurrent.py — directly launches the triton kernel -# without a torch.autograd.Function wrapper. Used by the MTP spec-decode +# without a torch.autograd.Function wrapper. Used by the MTP decode # verify path of the GDN (Gated DeltaNet) layer in Qwen3Next. # # Upstream source: flash-linear-attention / fused-recurrent gated delta rule. diff --git a/lightllm/common/basemodel/triton_kernel/linear_att/spec_state_params.py b/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py similarity index 94% rename from lightllm/common/basemodel/triton_kernel/linear_att/spec_state_params.py rename to lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py index 623e1bfc08..b81bf9ec16 100644 --- a/lightllm/common/basemodel/triton_kernel/linear_att/spec_state_params.py +++ b/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py @@ -4,7 +4,7 @@ @triton.jit -def _build_dynamic_spec_linear_att_state_params_kernel( +def _build_dynamic_mtp_linear_att_state_params_kernel( b_req_idx, b_mtp_index, req_to_mtp_state_index, @@ -47,13 +47,13 @@ def _build_dynamic_spec_linear_att_state_params_kernel( tl.store(out_num_accepted_tokens + sequence_index, accepted_state_index + 1, mask=is_start) -def build_dynamic_spec_linear_att_state_params( +def build_dynamic_mtp_linear_att_state_params( b_req_idx: torch.Tensor, b_mtp_index: torch.Tensor, req_to_mtp_state_index: torch.Tensor, hold_req_id: int, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Convert compact speculative-verify rows to variable-length GDN sequences. + """Convert compact MTP verify rows to variable-length GDN sequences. Rows remain request-major after dynamic verification trimming, but each request can contribute a different number of rows. ``b_mtp_index`` starts at zero for @@ -95,7 +95,7 @@ def build_dynamic_spec_linear_att_state_params( b_conv_buffer_idx = torch.empty_like(b_req_idx) b_num_accepted_tokens = torch.empty_like(b_req_idx) - _build_dynamic_spec_linear_att_state_params_kernel[(1,)]( + _build_dynamic_mtp_linear_att_state_params_kernel[(1,)]( b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, req_to_mtp_state_index=req_to_mtp_state_index, diff --git a/unit_tests/common/basemodel/attention/linear/test_gdn.py b/unit_tests/common/basemodel/attention/linear/test_gdn.py index e01b1a04ce..12b6996a83 100644 --- a/unit_tests/common/basemodel/attention/linear/test_gdn.py +++ b/unit_tests/common/basemodel/attention/linear/test_gdn.py @@ -60,6 +60,101 @@ def test_prefill_casts_final_state_to_cache_dtype(monkeypatch, cache_dtype): assert torch.equal(ssm_states, final_state.to(cache_dtype)) +def _create_decode_state(*, draft_step, dynamic_layout, b_req_idx, req_to_mtp_state_index=None): + model = SimpleNamespace( + is_mtp_draft_model=False, + mtp_manager=SimpleNamespace(get_decode_draft_step=lambda _: draft_step), + ) + infer_state = SimpleNamespace( + batch_size=b_req_idx.shape[0], + b_req_idx=b_req_idx, + b_mtp_index=torch.zeros_like(b_req_idx), + req_manager=SimpleNamespace( + HOLD_REQUEST_ID=req_to_mtp_state_index.shape[0] - 1 if req_to_mtp_state_index is not None else -1, + req_to_mtp_state_index=req_to_mtp_state_index, + ), + ) + return gdn.LinearAttDecodeAttState( + backend=SimpleNamespace( + model=model, + uses_dynamic_spec_verify_layout=lambda: dynamic_layout, + ), + infer_state=infer_state, + ) + + +def test_decode_state_initializes_normal_layout(): + b_req_idx = torch.tensor([2, 4], dtype=torch.int32) + state = _create_decode_state(draft_step=0, dynamic_layout=False, b_req_idx=b_req_idx) + + state.init_state() + + assert state.b_conv_buffer_idx is b_req_idx + assert state.b_ssm_buffer_idx is b_req_idx + assert state.b1_mtp_cu_q_seq_len is None + assert state.b_num_accepted_tokens is None + + +def test_decode_state_initializes_fixed_mtp_layout(): + b_req_idx = torch.tensor([2, 2, 2, 4, 4, 4], dtype=torch.int32) + req_to_mtp_state_index = torch.tensor([0, 0, 1, 0, 2, 0], dtype=torch.int32) + state = _create_decode_state( + draft_step=2, + dynamic_layout=False, + b_req_idx=b_req_idx, + req_to_mtp_state_index=req_to_mtp_state_index, + ) + + state.init_state() + + torch.testing.assert_close(state.b1_mtp_cu_q_seq_len, torch.tensor([0, 3, 6], dtype=torch.int32)) + torch.testing.assert_close(state.b_conv_buffer_idx, torch.tensor([2, 4], dtype=torch.int32)) + torch.testing.assert_close(state.b_num_accepted_tokens, torch.tensor([2, 3], dtype=torch.int32)) + torch.testing.assert_close( + state.b_ssm_buffer_idx, + torch.tensor([[6, 7, 8], [12, 13, 14]], dtype=torch.int32), + ) + + +def test_decode_state_initializes_dynamic_mtp_layout(monkeypatch): + b_req_idx = torch.tensor([2, 2, 4, 4, 4], dtype=torch.int32) + req_to_mtp_state_index = torch.tensor([0, 0, 1, 0, 2, 0], dtype=torch.int32) + expected_cu_q_seq_len = torch.tensor([0, 2, 5, 5, 5, 5], dtype=torch.int32) + expected_conv_buffer_idx = torch.tensor([2, 4, 5, 5, 5], dtype=torch.int32) + expected_num_accepted_tokens = torch.tensor([2, 3, 1, 1, 1], dtype=torch.int32) + build_calls = [] + + def build_params(**kwargs): + build_calls.append(kwargs) + return expected_cu_q_seq_len, expected_conv_buffer_idx, expected_num_accepted_tokens + + monkeypatch.setattr(gdn, "build_dynamic_mtp_linear_att_state_params", build_params) + state = _create_decode_state( + draft_step=2, + dynamic_layout=True, + b_req_idx=b_req_idx, + req_to_mtp_state_index=req_to_mtp_state_index, + ) + + state.init_state() + + assert state.b1_mtp_cu_q_seq_len is expected_cu_q_seq_len + assert state.b_conv_buffer_idx is expected_conv_buffer_idx + assert state.b_num_accepted_tokens is expected_num_accepted_tokens + torch.testing.assert_close( + state.b_ssm_buffer_idx, + torch.tensor( + [[6, 7, 8], [12, 13, 14], [15, 16, 17], [15, 16, 17], [15, 16, 17]], + dtype=torch.int32, + ), + ) + assert len(build_calls) == 1 + assert build_calls[0]["b_req_idx"] is state.infer_state.b_req_idx + assert build_calls[0]["b_mtp_index"] is state.infer_state.b_mtp_index + assert build_calls[0]["req_to_mtp_state_index"] is req_to_mtp_state_index + assert build_calls[0]["hold_req_id"] == 5 + + def test_linear_backend_is_abstract_and_triton_supplies_chunk_kernel(): assert inspect.isabstract(gdn.LinearAttBackend) backend = object.__new__(TritonLinearAttBackend) diff --git a/unit_tests/common/basemodel/triton_kernel/linear_att/test_causal_conv1d_spec.py b/unit_tests/common/basemodel/triton_kernel/linear_att/test_causal_conv1d_mtp.py similarity index 98% rename from unit_tests/common/basemodel/triton_kernel/linear_att/test_causal_conv1d_spec.py rename to unit_tests/common/basemodel/triton_kernel/linear_att/test_causal_conv1d_mtp.py index 3ac59a41ed..f7c742db01 100644 --- a/unit_tests/common/basemodel/triton_kernel/linear_att/test_causal_conv1d_spec.py +++ b/unit_tests/common/basemodel/triton_kernel/linear_att/test_causal_conv1d_mtp.py @@ -3,7 +3,7 @@ import pytest import torch -from lightllm.common.basemodel.triton_kernel.linear_att.causal_conv1d_spec import causal_conv1d_update +from lightllm.common.basemodel.triton_kernel.linear_att.causal_conv1d_mtp import causal_conv1d_update def causal_conv1d_ref( @@ -661,9 +661,9 @@ def test_multi_step_decode_sliding_window(width, num_steps): @pytest.mark.parametrize("width", [2, 3, 4]) @pytest.mark.parametrize("mtp_step", [1, 2, 3]) -def test_spec_decode_multi_token_per_step(width, mtp_step): +def test_mtp_decode_multi_token_per_step(width, mtp_step): """ - Spec-decode: each step processes (mtp_step+1) tokens. After each step, + MTP decode: each step processes (mtp_step+1) tokens. After each step, some tokens are accepted (num_accepted_tokens varies). The conv_state sliding window must correctly preserve history across steps. """ @@ -683,7 +683,7 @@ def test_spec_decode_multi_token_per_step(width, mtp_step): idxs = torch.zeros(batch, device=device, dtype=torch.int32) qsl = torch.tensor([0, seqlen], device=device, dtype=torch.int32) - # Phase 1: one-step spec decode triton + # Phase 1: one-step MTP decode triton x_full = torch.randn(seqlen, dim, device=device, dtype=torch.float32) out_triton = causal_conv1d_update( x_full.clone().half(), @@ -715,11 +715,11 @@ def test_spec_decode_multi_token_per_step(width, mtp_step): rtol, atol = 1e-2, 1e-2 assert torch.allclose(out_triton, ref_out, rtol=rtol, atol=atol), ( - f"Spec-decode mismatch: width={width}, mtp_step={mtp_step}\n" + f"MTP decode mismatch: width={width}, mtp_step={mtp_step}\n" f"max diff={torch.abs(out_triton - ref_out).max().item():.6f}" ) - # Phase 2: multi-step spec decode with varying acceptance + # Phase 2: multi-step MTP decode with varying acceptance conv_state_step = conv_state.clone().half() for step in range(num_steps): @@ -758,7 +758,7 @@ def test_spec_decode_multi_token_per_step(width, mtp_step): ).float() assert torch.allclose(step_out, ref_step, rtol=rtol, atol=atol), ( - f"Multi-step spec step={step} mismatch: width={width}, mtp_step={mtp_step}\n" + f"Multi-step MTP step={step} mismatch: width={width}, mtp_step={mtp_step}\n" f"max diff={torch.abs(step_out - ref_step).max().item():.6f}" ) From 552b02af6d56884d878c49f17275b076d67d4529 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 08:31:28 +0000 Subject: [PATCH 021/103] refactor: clarify MTP SSM buffer shape --- lightllm/common/basemodel/attention/linear/gdn.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index 872920d779..ca6ceaec43 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -252,13 +252,15 @@ def _init_fixed_mtp_decode_state(self, draft_step: int): self._init_mtp_ssm_buffer_idx(mtp_size) def _init_mtp_ssm_buffer_idx(self, mtp_size: int): - # Each request owns one recurrent-state slot per verify row. + att_batch_size = self.b_conv_buffer_idx.shape[0] + # Each request owns mtp_size consecutive recurrent-state slots. + b_ssm_buffer_start_idx = (self.b_conv_buffer_idx * mtp_size).view(att_batch_size, 1) state_offsets = torch.arange( mtp_size, device=self.infer_state.b_req_idx.device, dtype=self.infer_state.b_req_idx.dtype, - ) - self.b_ssm_buffer_idx = self.b_conv_buffer_idx[:, None] * mtp_size + state_offsets[None, :] + ).view(1, mtp_size) + self.b_ssm_buffer_idx = b_ssm_buffer_start_idx + state_offsets # [att_batch_size, mtp_size] def decode_att( self, From a423012b5d8d52e8ff23e5f326c54d1c1f910ed7 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 09:01:23 +0000 Subject: [PATCH 022/103] refactor: share GPU attention workspaces --- .../common/basemodel/attention/base_att.py | 22 ++++++++ .../basemodel/attention/flashinfer/fp.py | 14 ++++- .../basemodel/attention/flashinfer/mla.py | 15 ++++- .../attention/flashinfer/test_workspace.py | 56 +++++++++++++++++++ 4 files changed, 101 insertions(+), 6 deletions(-) create mode 100644 unit_tests/common/basemodel/attention/flashinfer/test_workspace.py diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index 6dd8d2c368..a7e2d8122a 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -1,8 +1,11 @@ +import threading + import torch from abc import ABC, abstractmethod from dataclasses import dataclass from typing import Optional, TYPE_CHECKING, Tuple, Union, Dict +from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.envs_utils import get_env_start_args if TYPE_CHECKING: @@ -17,6 +20,8 @@ class BaseAttBackend: """ _instances = {} + _workspace_buffers = {} + _workspace_buffer_lock = threading.Lock() def __new__(cls, *args, **kwargs): """ @@ -33,6 +38,23 @@ def __new__(cls, *args, **kwargs): def __init__(self, model: "TpPartBaseModel"): self.model = model + @staticmethod + def get_gpu_workspace_buffer(key_name: str, workspace_size: int, dtype: torch.dtype = torch.int8) -> torch.Tensor: + """Return a process-local workspace shared by key name and CUDA device.""" + if not key_name: + raise ValueError("workspace key_name must not be empty") + if workspace_size <= 0: + raise ValueError(f"workspace_size must be positive, got {workspace_size}") + + device_id = get_current_device_id() + buffer_key = (device_id, key_name, workspace_size, dtype) + with BaseAttBackend._workspace_buffer_lock: + workspace_buffer = BaseAttBackend._workspace_buffers.get(buffer_key) + if workspace_buffer is None: + workspace_buffer = torch.empty(workspace_size, dtype=dtype, device=device_id) + BaseAttBackend._workspace_buffers[buffer_key] = workspace_buffer + return workspace_buffer + def create_att_prefill_state(self) -> "BasePrefillAttState": raise NotImplementedError("not impl") diff --git a/lightllm/common/basemodel/attention/flashinfer/fp.py b/lightllm/common/basemodel/attention/flashinfer/fp.py index e0c44e8ed5..ec0b42dca3 100644 --- a/lightllm/common/basemodel/attention/flashinfer/fp.py +++ b/lightllm/common/basemodel/attention/flashinfer/fp.py @@ -8,6 +8,9 @@ class FlashInferAttBackend(BaseAttBackend): + workspace_buffer_key = "flashinfer_fp" + workspace_buffer_size = 512 * 1024 * 1024 + def __init__(self, model): set_flashinfer_envs() super().__init__(model=model) @@ -16,7 +19,6 @@ def __init__(self, model): self.tp_kv_head_num = max(model.config["num_key_value_heads"] // tp_world_size, 1) head_dim = model.config["hidden_size"] // model.config["num_attention_heads"] self.head_dim = model.config.get("head_dim", head_dim) - self.workspace_buffer = torch.empty(512 * 1024 * 1024, dtype=torch.int8, device=get_current_device_id()) self.max_seq_length = model.max_seq_length self.kv_indices_buffer = [ torch.empty( @@ -65,7 +67,10 @@ def init_state(self): kv_indices, ) self.prefill_wrapper = flashinfer.prefill.BatchPrefillWithPagedKVCacheWrapper( - self.backend.workspace_buffer, + self.backend.get_gpu_workspace_buffer( + key_name=self.backend.workspace_buffer_key, + workspace_size=self.backend.workspace_buffer_size, + ), qo_indptr_buf=q_starts, paged_kv_indptr_buf=kv_starts, paged_kv_indices_buf=kv_indices, @@ -166,7 +171,10 @@ def init_state(self): assert self.decode_wrapper is None self.decode_wrapper = flashinfer.decode.BatchDecodeWithPagedKVCacheWrapper( - self.backend.workspace_buffer, + self.backend.get_gpu_workspace_buffer( + key_name=self.backend.workspace_buffer_key, + workspace_size=self.backend.workspace_buffer_size, + ), "NHD", use_cuda_graph=True, use_tensor_cores=True, diff --git a/lightllm/common/basemodel/attention/flashinfer/mla.py b/lightllm/common/basemodel/attention/flashinfer/mla.py index a71dd4d464..2da8b4423c 100644 --- a/lightllm/common/basemodel/attention/flashinfer/mla.py +++ b/lightllm/common/basemodel/attention/flashinfer/mla.py @@ -10,6 +10,9 @@ class MlaFlashInferAttBackend(BaseAttBackend): + workspace_buffer_key = "flashinfer_mla" + workspace_buffer_size = 256 * 1024 * 1024 + def __init__(self, model): set_flashinfer_envs() super().__init__(model=model) @@ -21,7 +24,6 @@ def __init__(self, model): self.v_head_dim = model.v_head_dim self.q_data_type = model.data_type self.kv_data_type = model.data_type - self.workspace_buffer = torch.empty(256 * 1024 * 1024, dtype=torch.int8, device=get_current_device_id()) self.max_seq_length = model.max_seq_length self.softmax_scale = (self.qk_nope_head_dim + self.qk_rope_head_dim) ** (-0.5) self.kv_indices_buffer = [ @@ -64,7 +66,11 @@ def init_state(self): kv_starts = self.infer_state.b1_cu_kv_seq_len.int() if self.prefill_wrapper is None: self.prefill_wrapper = flashinfer.prefill.BatchPrefillWithRaggedKVCacheWrapper( - self.backend.workspace_buffer, "NHD" + self.backend.get_gpu_workspace_buffer( + key_name=self.backend.workspace_buffer_key, + workspace_size=self.backend.workspace_buffer_size, + ), + "NHD", ) self.prefill_wrapper.plan( qo_indptr=q_starts, @@ -159,7 +165,10 @@ def init_state(self): assert self.decode_wrapper is None self.decode_wrapper = flashinfer.mla.BatchMLAPagedAttentionWrapper( - self.backend.workspace_buffer, + self.backend.get_gpu_workspace_buffer( + key_name=self.backend.workspace_buffer_key, + workspace_size=self.backend.workspace_buffer_size, + ), use_cuda_graph=True, qo_indptr=self.q_indptr, kv_indices=self.kv_indices, diff --git a/unit_tests/common/basemodel/attention/flashinfer/test_workspace.py b/unit_tests/common/basemodel/attention/flashinfer/test_workspace.py new file mode 100644 index 0000000000..924dead570 --- /dev/null +++ b/unit_tests/common/basemodel/attention/flashinfer/test_workspace.py @@ -0,0 +1,56 @@ +from lightllm.common.basemodel.attention import base_att +from lightllm.common.basemodel.attention.base_att import BaseAttBackend +from lightllm.common.basemodel.attention.flashinfer.fp import FlashInferAttBackend +from lightllm.common.basemodel.attention.flashinfer.fp8 import Fp8FlashInferAttBackend +from lightllm.common.basemodel.attention.flashinfer.mla import MlaFlashInferAttBackend + + +def test_workspace_is_shared_by_backend_family_and_device(monkeypatch): + allocations = [] + current_device_id = 0 + + def fake_empty(size, dtype, device): + workspace = object() + allocations.append((workspace, size, dtype, device)) + return workspace + + monkeypatch.setattr(base_att, "get_current_device_id", lambda: current_device_id) + monkeypatch.setattr(base_att.torch, "empty", fake_empty) + monkeypatch.setattr(BaseAttBackend, "_workspace_buffers", {}) + + def get_workspace(backend_cls): + return backend_cls.get_gpu_workspace_buffer( + key_name=backend_cls.workspace_buffer_key, + workspace_size=backend_cls.workspace_buffer_size, + ) + + fp_workspace = get_workspace(FlashInferAttBackend) + fp8_workspace = get_workspace(Fp8FlashInferAttBackend) + mla_workspace = get_workspace(MlaFlashInferAttBackend) + + assert fp8_workspace is fp_workspace + assert mla_workspace is not fp_workspace + assert [allocation[1] for allocation in allocations] == [ + 512 * 1024 * 1024, + 256 * 1024 * 1024, + ] + + current_device_id = 1 + other_device_workspace = get_workspace(FlashInferAttBackend) + + assert other_device_workspace is not fp_workspace + assert allocations[-1][3] == 1 + + +def test_workspace_key_includes_size_and_dtype(monkeypatch): + monkeypatch.setattr(base_att, "get_current_device_id", lambda: 0) + monkeypatch.setattr(base_att.torch, "empty", lambda *args, **kwargs: object()) + monkeypatch.setattr(BaseAttBackend, "_workspace_buffers", {}) + + small_workspace = BaseAttBackend.get_gpu_workspace_buffer("shared", 1024) + large_workspace = BaseAttBackend.get_gpu_workspace_buffer("shared", 2048) + fp16_workspace = BaseAttBackend.get_gpu_workspace_buffer("shared", 1024, dtype=base_att.torch.float16) + + assert BaseAttBackend.get_gpu_workspace_buffer("shared", 1024) is small_workspace + assert large_workspace is not small_workspace + assert fp16_workspace is not small_workspace From 4e18ceed7e0fc93bd0a488bb3cbcc8a7ceb605d5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 09:12:42 +0000 Subject: [PATCH 023/103] fix --- lightllm/common/basemodel/triton_kernel/norm/qk_norm.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/norm/qk_norm.py b/lightllm/common/basemodel/triton_kernel/norm/qk_norm.py index e09d63f752..e152a8dd83 100644 --- a/lightllm/common/basemodel/triton_kernel/norm/qk_norm.py +++ b/lightllm/common/basemodel/triton_kernel/norm/qk_norm.py @@ -27,8 +27,8 @@ def _rms_norm_fwd_fused( var = tl.sum(x * x, axis=0) / head_dim rstd = 1 / tl.sqrt(var + eps) # Normalize and apply linear transformation - w = tl.load(W + tl.arange(0, BLOCK_SIZE)) - x_hat = (x * rstd).to(X.dtype.element_ty) + w = tl.load(W + tl.arange(0, BLOCK_SIZE)).to(tl.float32) + x_hat = x * rstd y = x_hat * w # Write output tl.store(X + cols, y.to(X.dtype.element_ty)) From 027ce10e9d542717fcc98fb671aa6cc279723123 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 23:26:59 +0000 Subject: [PATCH 024/103] refactor: move FA3 utility test to file end --- .../basemodel/triton_kernel/fa3_utils.py | 66 +++++++++---------- 1 file changed, 33 insertions(+), 33 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/fa3_utils.py b/lightllm/common/basemodel/triton_kernel/fa3_utils.py index 154ca9637a..520e8a7859 100644 --- a/lightllm/common/basemodel/triton_kernel/fa3_utils.py +++ b/lightllm/common/basemodel/triton_kernel/fa3_utils.py @@ -62,39 +62,6 @@ def page_table_copy( ) -def test_page_table_copy(): - import torch - - batch_size, seq_len = 2, 8 - - req_to_token_indexs = torch.arange(batch_size * seq_len, dtype=torch.int32).reshape(batch_size, seq_len).cuda() - - page_table = torch.full((batch_size, seq_len), -1, dtype=torch.int32, device="cuda") - - b_req_idx = torch.tensor([0, 2, 1, 3], dtype=torch.int32, device="cuda")[::2] - print(b_req_idx.stride()) - - page_table_copy(page_table, req_to_token_indexs, b_req_idx) - - print("req_to_token_indexs:") - print(req_to_token_indexs.cpu().numpy()) - print("b_req_idx:", b_req_idx.cpu().numpy()) - print("page_table:") - print(page_table.cpu().numpy()) - - for batch in range(batch_size): - src_idx = b_req_idx[batch].item() - expected = req_to_token_indexs[src_idx].cpu().numpy() - got = page_table[batch].cpu().numpy() - assert (expected == got).all(), f"Batch {batch} mismatch: expected {expected}, got {got}" - - print("✅ Test passed!") - - -if __name__ == "__main__": - test_page_table_copy() - - @triton.jit def _build_dynamic_spec_fa3_decode_params_kernel( b_req_idx, @@ -286,3 +253,36 @@ def build_dynamic_spec_fa3_decode_params( num_stages=1, ) return b_q_seq_len, b_kv_seq_len, b_att_req_idx, b_att_seq_len + + +def test_page_table_copy(): + import torch + + batch_size, seq_len = 2, 8 + + req_to_token_indexs = torch.arange(batch_size * seq_len, dtype=torch.int32).reshape(batch_size, seq_len).cuda() + + page_table = torch.full((batch_size, seq_len), -1, dtype=torch.int32, device="cuda") + + b_req_idx = torch.tensor([0, 2, 1, 3], dtype=torch.int32, device="cuda")[::2] + print(b_req_idx.stride()) + + page_table_copy(page_table, req_to_token_indexs, b_req_idx) + + print("req_to_token_indexs:") + print(req_to_token_indexs.cpu().numpy()) + print("b_req_idx:", b_req_idx.cpu().numpy()) + print("page_table:") + print(page_table.cpu().numpy()) + + for batch in range(batch_size): + src_idx = b_req_idx[batch].item() + expected = req_to_token_indexs[src_idx].cpu().numpy() + got = page_table[batch].cpu().numpy() + assert (expected == got).all(), f"Batch {batch} mismatch: expected {expected}, got {got}" + + print("✅ Test passed!") + + +if __name__ == "__main__": + test_page_table_copy() From bf339c1a6bf0812f2227237b150eb10b5cea91d8 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Mon, 17 Aug 2026 23:39:21 +0000 Subject: [PATCH 025/103] refactor: align MTP utility naming --- ...mic_spec_utils.py => dynamic_mtp_utils.py} | 2 +- .../basemodel/triton_kernel/mtp_utils.py | 40 +++++++-------- .../mode_backend/chunked_prefill/impl.py | 2 +- .../mode_backend/dp_backend/impl.py | 50 +++++++++---------- .../router/model_infer/speculative/engine.py | 12 ++--- ...pec_utils.py => test_dynamic_mtp_utils.py} | 14 +++--- .../basemodel/triton_kernel/test_mtp_utils.py | 12 ++--- .../mode_backend/test_dp_spec_engine.py | 4 +- 8 files changed, 68 insertions(+), 68 deletions(-) rename lightllm/common/basemodel/triton_kernel/{dynamic_spec_utils.py => dynamic_mtp_utils.py} (98%) rename unit_tests/common/basemodel/triton_kernel/{test_dynamic_spec_utils.py => test_dynamic_mtp_utils.py} (96%) diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py similarity index 98% rename from lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py rename to lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py index 74bc9462f1..b935e7c8a7 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_spec_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -34,7 +34,7 @@ def _fwd_kernel_cumprod_scores( return -def sample_dynamic_spec_row_mask( +def sample_dynamic_mtp_row_mask( dynamic_batch_size: int, b_req_idx: torch.Tensor, req_to_next_token_scores: torch.Tensor, diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index 90cc4ec4be..7f05a975c4 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -5,7 +5,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.basemodel.mtp_manager import MtpManager -from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import sample_dynamic_spec_row_mask +from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import sample_dynamic_mtp_row_mask from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size @@ -14,7 +14,7 @@ def _fwd_kernel_mtp_verify( req_to_next_token_ids, req_to_next_token_ids_stride, new_next_token_ids, - spec_accept_len, + mtp_accept_len, b_req_mtp_start_loc, b_req_idx, accepted_index, @@ -43,7 +43,7 @@ def _fwd_kernel_mtp_verify( mismatch_positions = tl.where(match_mask, BLOCK_SIZE, offset) first_mismatch_pos = tl.min(mismatch_positions) accept_len = first_mismatch_pos + 1 - tl.store(spec_accept_len + cur_index, accept_len) + tl.store(mtp_accept_len + cur_index, accept_len) accpeted_index = tl.where((offset < accept_len), 1, 0) tl.store(accepted_index + req_offset, accpeted_index, mask=offset < req_mtp_num) return @@ -63,7 +63,7 @@ def mtp_verify( new_next_token_ids: (batch_size,) b_req_idx: (batch_size,) Returns: - spec_accept_len: (num_reqs,) + mtp_accept_len: (num_reqs,) accepted_index: (batch_size,) accepted_index: [1, 0, 1, 1, 0], 0 means the token is not accepted, 1 means the token is accepted. """ @@ -72,7 +72,7 @@ def mtp_verify( assert verify_width <= BLOCK_SIZE, f"verify_width must be less than {BLOCK_SIZE}" num_reqs = b_req_mtp_start_loc.shape[0] req_mtp_all_num = b_req_idx.shape[0] - spec_accept_len = torch.empty((num_reqs,), dtype=torch.int32, device=req_to_next_token_ids.device) + mtp_accept_len = torch.empty((num_reqs,), dtype=torch.int32, device=req_to_next_token_ids.device) accepted_index = torch.empty((req_mtp_all_num,), dtype=torch.int32, device=req_to_next_token_ids.device) grid = (num_reqs,) @@ -81,7 +81,7 @@ def mtp_verify( req_to_next_token_ids=req_to_next_token_ids, req_to_next_token_ids_stride=req_to_next_token_ids.stride(0), new_next_token_ids=new_next_token_ids, - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, accepted_index=accepted_index, @@ -90,7 +90,7 @@ def mtp_verify( num_warps=num_warps, num_stages=1, ) - return spec_accept_len, accepted_index + return mtp_accept_len, accepted_index @triton.jit @@ -103,7 +103,7 @@ def _fwd_kernel_mtp_scatter_next_token_ids( req_to_next_token_scores_stride, schedule_scores, schedule_scores_stride, - spec_accept_len, + mtp_accept_len, b_req_mtp_start_loc, b_req_idx, proposal_width, @@ -114,7 +114,7 @@ def _fwd_kernel_mtp_scatter_next_token_ids( cur_index = tl.program_id(0) req_start_loc = tl.load(b_req_mtp_start_loc + cur_index) - accept_len = tl.load(spec_accept_len + cur_index) + accept_len = tl.load(mtp_accept_len + cur_index) cur_req_idx = tl.load(b_req_idx + req_start_loc) offset = tl.arange(0, BLOCK_SIZE) selected_row = req_start_loc + accept_len - 1 @@ -152,7 +152,7 @@ def mtp_scatter_next_token_ids( b_req_mtp_start_loc: torch.Tensor, all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, - spec_accept_len: torch.Tensor, + mtp_accept_len: torch.Tensor, req_to_next_token_scores: Optional[torch.Tensor] = None, schedule_scores: Optional[torch.Tensor] = None, ): @@ -189,7 +189,7 @@ def mtp_scatter_next_token_ids( req_to_next_token_scores_stride=req_to_next_token_scores_stride, schedule_scores=schedule_scores_arg, schedule_scores_stride=schedule_scores_stride, - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, proposal_width=proposal_width, @@ -202,7 +202,7 @@ def mtp_scatter_next_token_ids( @triton.jit -def _fwd_kernel_compact_dynamic_spec_model_input( +def _fwd_kernel_compact_dynamic_mtp_model_input( input_ids, out_input_ids, b_req_idx, @@ -440,7 +440,7 @@ def _compact_decode_model_input( dummy_1d = model_input.b_req_idx BLOCK_SIZE = triton.next_power_of_2(old_batch_size) grid = (1,) - _fwd_kernel_compact_dynamic_spec_model_input[grid]( + _fwd_kernel_compact_dynamic_mtp_model_input[grid]( input_ids=model_input.input_ids if model_input.input_ids is not None else dummy_1d, out_input_ids=out_input_ids if out_input_ids is not None else dummy_1d, b_req_idx=model_input.b_req_idx, @@ -492,7 +492,7 @@ def _compact_decode_model_input( return model_input -def prepare_dynamic_spec_model_input( +def prepare_dynamic_mtp_model_input( model_input: ModelInput, req_num: int, dynamic_batch_size: int, @@ -501,7 +501,7 @@ def prepare_dynamic_spec_model_input( ): req_num = int(req_num) dynamic_batch_size = int(dynamic_batch_size) - assert not model_input.is_prefill, "prepare_dynamic_spec_model_input only supports decode inputs" + assert not model_input.is_prefill, "prepare_dynamic_mtp_model_input only supports decode inputs" assert req_to_next_token_scores is not None assert dynamic_batch_size >= req_num assert dynamic_batch_size <= model_input.batch_size @@ -514,7 +514,7 @@ def prepare_dynamic_spec_model_input( # All compaction work stays on the current CUDA stream and needs no host sync. model_input.to_cuda() - selected_row_mask = sample_dynamic_spec_row_mask( + selected_row_mask = sample_dynamic_mtp_row_mask( dynamic_batch_size=dynamic_batch_size, b_req_idx=model_input.b_req_idx, req_to_next_token_scores=req_to_next_token_scores, @@ -607,7 +607,7 @@ def _fwd_kernel_linear_att_mtp_state_index_update( return -def linear_att_spec_state_index_update( +def linear_att_mtp_state_index_update( req_to_mtp_state_index: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, b_req_idx: torch.Tensor, @@ -657,13 +657,13 @@ def test_mtp_verify(): all_next_token_ids = torch.tensor( [[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], [13, 14, 15]], dtype=torch.int64, device="cuda" ) - spec_accept_len, accepted_index = mtp_verify( + mtp_accept_len, accepted_index = mtp_verify( req_to_next_token_ids, b_req_mtp_start_loc, new_next_token_ids, b_req_idx ) mtp_scatter_next_token_ids( - req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, spec_accept_len + req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, mtp_accept_len ) - print(spec_accept_len) + print(mtp_accept_len) print(req_to_next_token_ids) print(accepted_index) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 9476d98ecb..22bcee68fc 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -306,7 +306,7 @@ def decode_mtp( b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=proposal.token_ids, b_req_idx=model_input.b_req_idx, - spec_accept_len=mtp_accept_len, + mtp_accept_len=mtp_accept_len, schedule_scores=proposal.schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 04ab0de46d..054e84a6c2 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -504,7 +504,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output = self.model.forward(model_input) - spec_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None + mtp_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None if req_num > 0: logits = model_output.logits[0:req_num, :] b_mtp_index_cpu = b_mtp_index_cpu[0:req_num] @@ -528,7 +528,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): dtype=torch.int32, ).cuda(non_blocking=True) - spec_accept_len, accepted_index = self.spec_engine.verify_tokens( + mtp_accept_len, accepted_index = self.spec_engine.verify_tokens( next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, @@ -538,9 +538,9 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): key="accepted_index", gpu_tensor=accepted_index, ) - spec_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="spec_accept_len", - gpu_tensor=spec_accept_len, + mtp_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="mtp_accept_len", + gpu_tensor=mtp_accept_len, ) verify_event = torch.cuda.Event() @@ -551,7 +551,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): model_output=model_output, next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, req_num=req_num, ) if req_num > 0: @@ -570,7 +570,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): verify_event.synchronize() self.spec_engine.record_request_spec_metrics( decode_reqs=decode_reqs, - accept_lengths_cpu=spec_accept_len_cpu, + accept_lengths_cpu=mtp_accept_len_cpu, ) verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) @@ -608,7 +608,7 @@ def _draft_decode_vanilla( model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, - spec_accept_len: torch.Tensor, + mtp_accept_len: torch.Tensor, req_num: int, ): all_next_token_ids = [] @@ -640,7 +640,7 @@ def _draft_decode_vanilla( all_next_token_ids=stacked_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx[:req_num], - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, ) return None @@ -650,7 +650,7 @@ def _draft_decode_eagle( model_output: ModelOutput, next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, - spec_accept_len: torch.Tensor, + mtp_accept_len: torch.Tensor, req_num: int, ): verify_width = self.max_draft_step + 1 @@ -676,7 +676,7 @@ def _draft_decode_eagle( device=model_input.b_req_idx.device, ) if real_request_num > 0: - padded_accept_len[:real_request_num].copy_(spec_accept_len) + padded_accept_len[:real_request_num].copy_(mtp_accept_len) # DP keeps the target verify layout padded for collective shape # agreement. The proposer still follows the common topology: one @@ -696,7 +696,7 @@ def _draft_decode_eagle( b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=proposal.token_ids[:req_num], b_req_idx=model_input.b_req_idx[:req_num], - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, schedule_scores=(proposal.schedule_scores[:req_num] if proposal.schedule_scores is not None else None), ) return proposal.extra_mem_indexes_cpu @@ -817,7 +817,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf logits0 = model_output0.logits logits1 = model_output1.logits run_reqs = run_reqs0 + run_reqs1 - b_req_idx, spec_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None, None + b_req_idx, mtp_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None, None if (req_num0 + req_num1) > 0: logits = torch.empty( (req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device @@ -851,7 +851,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if self.is_linear_att_mixed_model else None ) - spec_accept_len, accepted_index = self.spec_engine.verify_tokens( + mtp_accept_len, accepted_index = self.spec_engine.verify_tokens( next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, @@ -861,9 +861,9 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf key="accepted_index", gpu_tensor=accepted_index, ) - spec_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( - key="spec_accept_len", - gpu_tensor=spec_accept_len, + mtp_accept_len_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( + key="mtp_accept_len", + gpu_tensor=mtp_accept_len, ) all_next_token_ids.append(next_token_ids) @@ -877,7 +877,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf model_output1=model_output1, b_req_idx=b_req_idx, next_token_ids=next_token_ids, - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, req_num0=req_num0, req_num1=req_num1, @@ -897,7 +897,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf verify_event.synchronize() self.spec_engine.record_request_spec_metrics( decode_reqs=decode_reqs, - accept_lengths_cpu=spec_accept_len_cpu, + accept_lengths_cpu=mtp_accept_len_cpu, ) verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) @@ -937,7 +937,7 @@ def _draft_decode_vanilla_overlap( model_output1: ModelOutput, b_req_idx: torch.Tensor, next_token_ids: torch.Tensor = None, - spec_accept_len: torch.Tensor = None, + mtp_accept_len: torch.Tensor = None, b_req_mtp_start_loc: torch.Tensor = None, req_num0: int = 0, req_num1: int = 0, @@ -990,7 +990,7 @@ def _draft_decode_vanilla_overlap( all_next_token_ids=stacked_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, ) return None @@ -1002,7 +1002,7 @@ def _draft_decode_eagle_overlap( model_output1: ModelOutput, b_req_idx: torch.Tensor, next_token_ids: torch.Tensor = None, - spec_accept_len: torch.Tensor = None, + mtp_accept_len: torch.Tensor = None, b_req_mtp_start_loc: torch.Tensor = None, req_num0: int = 0, req_num1: int = 0, @@ -1038,10 +1038,10 @@ def _draft_decode_eagle_overlap( device=model_input1.b_req_idx.device, ) if real_request_num0 > 0: - padded_accept_len0[:real_request_num0].copy_(spec_accept_len[:real_request_num0]) + padded_accept_len0[:real_request_num0].copy_(mtp_accept_len[:real_request_num0]) if real_request_num1 > 0: padded_accept_len1[:real_request_num1].copy_( - spec_accept_len[real_request_num0 : real_request_num0 + real_request_num1] + mtp_accept_len[real_request_num0 : real_request_num0 + real_request_num1] ) proposal = self.spec_engine.propose_next_overlap( @@ -1063,7 +1063,7 @@ def _draft_decode_eagle_overlap( b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=proposal.token_ids, b_req_idx=b_req_idx, - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, schedule_scores=proposal.schedule_scores, ) return proposal.extra_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/speculative/engine.py b/lightllm/server/router/model_infer/speculative/engine.py index c8a70f0f0a..e68b01fa34 100644 --- a/lightllm/server/router/model_infer/speculative/engine.py +++ b/lightllm/server/router/model_infer/speculative/engine.py @@ -7,7 +7,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.mtp_utils import ( - linear_att_spec_state_index_update, + linear_att_mtp_state_index_update, mtp_scatter_next_token_ids, mtp_verify, ) @@ -106,9 +106,9 @@ def prepare_decode_model_input( if plan.dynamic_batch_size == model_input.batch_size: return model_input, None - from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_spec_model_input + from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_mtp_model_input - model_input, selected_row_mask = prepare_dynamic_spec_model_input( + model_input, selected_row_mask = prepare_dynamic_mtp_model_input( model_input=model_input, req_num=req_num, dynamic_batch_size=plan.dynamic_batch_size, @@ -138,7 +138,7 @@ def verify_tokens( ) if self.backend.is_linear_att_mixed_model: assert b_mtp_index is not None - linear_att_spec_state_index_update( + linear_att_mtp_state_index_update( req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, @@ -221,7 +221,7 @@ def scatter_next_tokens( b_req_mtp_start_loc: torch.Tensor, all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, - spec_accept_len: torch.Tensor, + mtp_accept_len: torch.Tensor, schedule_scores: Optional[torch.Tensor] = None, ) -> None: mtp_scatter_next_token_ids( @@ -229,7 +229,7 @@ def scatter_next_tokens( b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=all_next_token_ids, b_req_idx=b_req_idx, - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, req_to_next_token_scores=( self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_scores if schedule_scores is not None diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py similarity index 96% rename from unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py rename to unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py index 917fc1cb04..88a2e4b999 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_dynamic_spec_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py @@ -3,9 +3,9 @@ import triton import numpy as np -from lightllm.common.basemodel.triton_kernel.dynamic_spec_utils import ( +from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import ( _fwd_kernel_cumprod_scores, - sample_dynamic_spec_row_mask, + sample_dynamic_mtp_row_mask, ) @@ -140,7 +140,7 @@ def test_sample_select_count(): ) all_num = req_num * (max_draft_step + 1) for dynamic_batch_size in [3, 8, all_num]: - select = sample_dynamic_spec_row_mask( + select = sample_dynamic_mtp_row_mask( dynamic_batch_size=dynamic_batch_size, b_req_idx=b_req_idx, req_to_next_token_scores=scores.clone(), @@ -163,7 +163,7 @@ def test_sample_accepts_numpy_scalar_dynamic_batch_size(): [1.0, 0.99, 0.99, 0.99], ], ) - select = sample_dynamic_spec_row_mask( + select = sample_dynamic_mtp_row_mask( dynamic_batch_size=np.int64(8), b_req_idx=b_req_idx, req_to_next_token_scores=scores, @@ -185,7 +185,7 @@ def test_sample_topk_by_cumprod_score(): ) flat_scores = _flat_cumprod_scores(b_req_idx, scores, max_draft_step) for dynamic_batch_size in [1, 4, 8, 12]: - select = sample_dynamic_spec_row_mask( + select = sample_dynamic_mtp_row_mask( dynamic_batch_size=dynamic_batch_size, b_req_idx=b_req_idx, req_to_next_token_scores=scores.clone(), @@ -205,7 +205,7 @@ def test_sample_picks_highest_cumprod_rows(): ], ) flat_scores = _flat_cumprod_scores(b_req_idx, scores, max_draft_step) - select = sample_dynamic_spec_row_mask( + select = sample_dynamic_mtp_row_mask( dynamic_batch_size=2, b_req_idx=b_req_idx, req_to_next_token_scores=scores.clone(), @@ -221,7 +221,7 @@ def test_sample_single_request(): max_draft_step = 2 scores, b_req_idx = _make_batch_scores(1, max_draft_step, rows=[[1.0, 0.5, 0.25]]) flat_scores = _flat_cumprod_scores(b_req_idx, scores, max_draft_step) - select = sample_dynamic_spec_row_mask( + select = sample_dynamic_mtp_row_mask( dynamic_batch_size=2, b_req_idx=b_req_idx, req_to_next_token_scores=scores.clone(), diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py index 12cd766d67..faefb2863c 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -18,7 +18,7 @@ def _reset_mtp_manager(): MtpManager._instance = None -def test_compact_dynamic_spec_model_input(monkeypatch): +def test_compact_dynamic_mtp_model_input(monkeypatch): monkeypatch.setenv( "LIGHTLLM_START_ARGS", json.dumps( @@ -64,7 +64,7 @@ def test_compact_dynamic_spec_model_input(monkeypatch): device="cuda", ) - compacted_input, selected_row_mask = mtp_utils.prepare_dynamic_spec_model_input( + compacted_input, selected_row_mask = mtp_utils.prepare_dynamic_mtp_model_input( model_input=model_input, req_num=3, dynamic_batch_size=8, @@ -164,7 +164,7 @@ def test_mtp_verify_scatter_and_start_locations(): device="cuda", ) - spec_accept_len, accepted_index = mtp_utils.mtp_verify( + mtp_accept_len, accepted_index = mtp_utils.mtp_verify( req_to_next_token_ids, b_req_mtp_start_loc, new_next_token_ids, b_req_idx ) mtp_utils.mtp_scatter_next_token_ids( @@ -172,14 +172,14 @@ def test_mtp_verify_scatter_and_start_locations(): b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=all_next_token_ids, b_req_idx=b_req_idx, - spec_accept_len=spec_accept_len, + mtp_accept_len=mtp_accept_len, req_to_next_token_scores=req_to_next_token_scores, schedule_scores=schedule_scores, ) torch.cuda.synchronize() assert torch.equal(b_req_mtp_start_loc.cpu(), torch.tensor([0, 2], dtype=torch.int64)) - assert torch.equal(spec_accept_len.cpu(), torch.tensor([1, 1], dtype=torch.int32)) + assert torch.equal(mtp_accept_len.cpu(), torch.tensor([1, 1], dtype=torch.int32)) assert torch.equal(accepted_index.cpu(), torch.tensor([1, 0, 1, 0, 0], dtype=torch.int32)) assert torch.equal( req_to_next_token_ids.cpu(), @@ -206,7 +206,7 @@ def test_mtp_scatter_handles_zero_draft_step(): b_req_mtp_start_loc=torch.tensor([0], dtype=torch.int32, device="cuda"), all_next_token_ids=torch.tensor([[42]], dtype=torch.int64, device="cuda"), b_req_idx=torch.tensor([0], dtype=torch.int32, device="cuda"), - spec_accept_len=torch.tensor([1], dtype=torch.int32, device="cuda"), + mtp_accept_len=torch.tensor([1], dtype=torch.int32, device="cuda"), req_to_next_token_scores=req_to_next_token_scores, schedule_scores=torch.empty((1, 0), dtype=torch.float32, device="cuda"), ) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py index e9118b71bd..0aa3f68031 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py @@ -62,7 +62,7 @@ def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(): model_output=model_output, next_token_ids=next_token_ids, b_req_mtp_start_loc=real_start_locs, - spec_accept_len=real_accept_len, + mtp_accept_len=real_accept_len, req_num=8, ) @@ -101,7 +101,7 @@ def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(): model_output1=model_output1, b_req_idx=b_req_idx, next_token_ids=next_token_ids, - spec_accept_len=accept_len, + mtp_accept_len=accept_len, b_req_mtp_start_loc=start_locs, req_num0=8, req_num1=16, From 35cc0ff88ac2047f75ae6047dfc93300a371e5a2 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 00:07:02 +0000 Subject: [PATCH 026/103] fix: clarify MRoPE positions for draft cache extension --- lightllm/common/basemodel/batch_objs.py | 8 +- lightllm/models/qwen2_vl/infer_struct.py | 46 +++++++++-- .../models/qwen2_vl/test_infer_struct.py | 82 +++++++++++++++++-- 3 files changed, 118 insertions(+), 18 deletions(-) diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index db901f5a6d..0bfbc634b1 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -43,7 +43,9 @@ class ModelInput: mem_indexes: torch.Tensor = None is_prefill: bool = False b_ready_cache_len: torch.Tensor = None - # MRoPE position offset; preserve it across row-aligned input transforms. + # Request/row-aligned MRoPE position offset. Normal prompt prefill leaves + # it unset; decode and one-token-per-row MTP draft KV commits carry it. + # Row-aligned input transforms must preserve the tensor unchanged. b_position_delta: torch.Tensor = None b_prefill_start_loc: torch.Tensor = None multimodal_params: list = None @@ -78,8 +80,10 @@ def to_cuda(self): self.b_ready_cache_len = self.b_ready_cache_len.cuda(non_blocking=True) if self.b_position_delta is not None: self.b_position_delta = self.b_position_delta.cuda(non_blocking=True) - assert self.is_prefill is False, "b_position_delta should only be used in decode phase." else: + # Decode always needs the request-level MRoPE delta. Prefill may + # omit it for a normal prompt, while MTP draft KV commit prefill + # deliberately carries it because its rows use decode positions. assert self.is_prefill is True, "decode ModelInput should provide b_position_delta." if self.b_prefill_start_loc is not None: diff --git a/lightllm/models/qwen2_vl/infer_struct.py b/lightllm/models/qwen2_vl/infer_struct.py index f3ae4ba668..399dfa866b 100644 --- a/lightllm/models/qwen2_vl/infer_struct.py +++ b/lightllm/models/qwen2_vl/infer_struct.py @@ -17,22 +17,54 @@ def init_some_extra_state(self, model): rope_scaling = model.config.get("rope_scaling", {}) self.rope_type = rope_scaling.get("rope_type", rope_scaling.get("type", None)) InferStateInfo.init_some_extra_state(self, model) + + # Case 1: normal prompt prefill. There is no request-level position + # delta yet, so build the complete 3-axis MRoPE positions from the + # prompt's image/video layout. if self.is_prefill and self.b_position_delta is None: self.position_ids = self.get_mrope_position(self.multimodal_params) + + # Case 2: normal decode. Base position_ids contains one scalar position + # per request; add the cached multimodal delta and broadcast the result + # to MRoPE's temporal/height/width axes. + elif not self.is_prefill: + assert self.b_position_delta is not None, "decode requires b_position_delta" + self._apply_mrope_position_delta() + + # Case 3: extend the draft-model KV cache after target verification. + # These rows originally came from target-model decode verification, + # potentially with multiple token rows belonging to the same request. + # The proposer reuses that row-aligned ModelInput and changes + # is_prefill to True so the draft model can replay all verified rows + # through its prefill kernel and write them into the draft KV cache in + # one forward. Therefore is_prefill describes the kernel/cache-write + # path here; it does not mean that this is the original prompt prefill. + # + # Every row is a post-prompt token and already carries the request-level + # b_position_delta derived from the original multimodal prompt. Its + # MRoPE position must consequently be calculated in the same way as a + # decode token: base position plus b_position_delta. Rebuilding MRoPE + # positions from multimodal_params would be incorrect because the + # original image/video token layout is no longer being prefetched and + # multimodal_params may contain only empty row-aligned placeholders. else: - b_position_delta = self.b_position_delta.to(dtype=self.position_ids.dtype) - if self.is_prefill: - # Speculative draft commits use one prefill row per verified - # token, whose decode-time MRoPE offset is already available. - assert b_position_delta.shape == self.position_ids.shape - position_ids = self.position_ids + b_position_delta - self.position_ids = position_ids.unsqueeze(0).expand(3, -1) + assert self.b_position_delta is not None + self._apply_mrope_position_delta() self.position_ids = self.position_ids.contiguous() self.position_cos = model._cos_cached[self.position_ids] self.position_sin = model._sin_cached[self.position_ids] return + def _apply_mrope_position_delta(self): + b_position_delta = self.b_position_delta.to(dtype=self.position_ids.dtype) + assert b_position_delta.shape == self.position_ids.shape, ( + "b_position_delta must align with position_ids, " + f"got delta_shape={b_position_delta.shape}, position_shape={self.position_ids.shape}" + ) + position_ids = self.position_ids + b_position_delta + self.position_ids = position_ids.unsqueeze(0).expand(3, -1) + def get_mrope_position(self, multimodal_params: List[dict]) -> torch.Tensor: if len(multimodal_params) == 0: return self.position_ids.unsqueeze(0).expand(3, -1) diff --git a/unit_tests/models/qwen2_vl/test_infer_struct.py b/unit_tests/models/qwen2_vl/test_infer_struct.py index e62b188940..e5aa40f4a8 100644 --- a/unit_tests/models/qwen2_vl/test_infer_struct.py +++ b/unit_tests/models/qwen2_vl/test_infer_struct.py @@ -1,29 +1,93 @@ from types import SimpleNamespace +import pytest import torch +from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.basemodel.infer_struct import InferStateInfo from lightllm.models.qwen2_vl.infer_struct import Qwen2VLInferStateInfo -def test_draft_commit_prefill_uses_existing_position_delta(monkeypatch): - def init_single_token_prefill(self, model): +def _patch_base_position_ids(monkeypatch): + def init_positions(self, model): self.position_ids = torch.tensor([9, 19], dtype=torch.int32) self.b_q_seq_len = torch.ones(2, dtype=torch.int32) - monkeypatch.setattr(InferStateInfo, "init_some_extra_state", init_single_token_prefill) + monkeypatch.setattr(InferStateInfo, "init_some_extra_state", init_positions) - infer_state = Qwen2VLInferStateInfo() - infer_state.is_prefill = True - infer_state.b_position_delta = torch.tensor([3, 5], dtype=torch.int32) - infer_state.multimodal_params = [{"images": [], "audios": []}] * 2 - model = SimpleNamespace( + +def _make_model(): + return SimpleNamespace( config={"rope_scaling": {}}, _cos_cached=torch.arange(32, dtype=torch.float32).view(32, 1), _sin_cached=torch.arange(32, dtype=torch.float32).view(32, 1), ) - infer_state.init_some_extra_state(model) + +def test_normal_prompt_prefill_builds_multimodal_positions(monkeypatch): + _patch_base_position_ids(monkeypatch) + expected_position_ids = torch.tensor([[1, 2], [3, 4], [5, 6]], dtype=torch.int32) + monkeypatch.setattr( + Qwen2VLInferStateInfo, + "get_mrope_position", + lambda self, multimodal_params: expected_position_ids, + ) + + infer_state = Qwen2VLInferStateInfo() + infer_state.is_prefill = True + infer_state.b_position_delta = None + infer_state.multimodal_params = [{"images": [], "audios": []}] * 2 + + infer_state.init_some_extra_state(_make_model()) + + assert torch.equal(infer_state.position_ids, expected_position_ids) + + +def test_normal_decode_applies_position_delta(monkeypatch): + _patch_base_position_ids(monkeypatch) + + infer_state = Qwen2VLInferStateInfo() + infer_state.is_prefill = False + infer_state.b_position_delta = torch.tensor([3, 5], dtype=torch.int32) + infer_state.multimodal_params = [{"images": [], "audios": []}] * 2 + + infer_state.init_some_extra_state(_make_model()) expected_position_ids = torch.tensor([[12, 24]] * 3, dtype=torch.int32) assert torch.equal(infer_state.position_ids, expected_position_ids) + + +def test_draft_commit_prefill_uses_existing_position_delta(monkeypatch): + _patch_base_position_ids(monkeypatch) + + infer_state = Qwen2VLInferStateInfo() + infer_state.is_prefill = True + infer_state.b_position_delta = torch.tensor([3, 5], dtype=torch.int32) + infer_state.multimodal_params = [{"images": [], "audios": []}] * 2 + + infer_state.init_some_extra_state(_make_model()) + + expected_position_ids = torch.tensor([[12, 24]] * 3, dtype=torch.int32) + assert torch.equal(infer_state.position_ids, expected_position_ids) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_draft_commit_prefill_model_input_accepts_position_delta(): + model_input = ModelInput( + batch_size=2, + total_token_num=2, + max_q_seq_len=1, + max_kv_seq_len=8, + input_ids=torch.tensor([1, 2], dtype=torch.int64), + b_req_idx=torch.tensor([0, 1], dtype=torch.int32), + b_mtp_index=torch.zeros(2, dtype=torch.int32), + b_seq_len=torch.tensor([4, 6], dtype=torch.int32), + b_position_delta=torch.tensor([3, 5], dtype=torch.int32), + mem_indexes_cpu=torch.tensor([10, 11], dtype=torch.int32), + is_prefill=True, + multimodal_params=[{"images": [], "audios": []}] * 2, + ) + + model_input.to_cuda() + + assert model_input.b_position_delta.is_cuda From 7fa888fe5637c2e9f57888e89e078fb165090b17 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 00:48:13 +0000 Subject: [PATCH 027/103] refactor: split draft model registry --- lightllm/models/__init__.py | 3 +- lightllm/models/deepseek_mtp/model.py | 2 +- lightllm/models/draft_registry.py | 46 ++++++++++++++++++++++ lightllm/models/glm4_moe_lite_mtp/model.py | 2 +- lightllm/models/mistral_mtp/model.py | 2 +- lightllm/models/qwen3_5_dflash/model.py | 2 +- lightllm/models/qwen3_5_dspark/model.py | 2 +- lightllm/models/qwen3_5_moe_mtp/model.py | 2 +- lightllm/models/qwen3_5_mtp/model.py | 2 +- lightllm/models/qwen3_dflash/model.py | 2 +- lightllm/models/qwen3_dspark/model.py | 2 +- lightllm/models/qwen3_eagle/model.py | 2 +- lightllm/models/qwen3_moe_mtp/model.py | 2 +- lightllm/models/registry.py | 42 +------------------- 14 files changed, 60 insertions(+), 53 deletions(-) create mode 100644 lightllm/models/draft_registry.py diff --git a/lightllm/models/__init__.py b/lightllm/models/__init__.py index c1063e5bb3..c7e9a59aad 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -54,4 +54,5 @@ from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel from lightllm.models.qwen3_eagle.model import Qwen3EagleModel from lightllm.models.qwen3_moe_mtp.model import Qwen3MOEMTPModel -from .registry import get_draft_model_class, get_model, get_model_class +from .draft_registry import get_draft_model_class +from .registry import get_model, get_model_class diff --git a/lightllm/models/deepseek_mtp/model.py b/lightllm/models/deepseek_mtp/model.py index 959528c1e9..6668678e6a 100644 --- a/lightllm/models/deepseek_mtp/model.py +++ b/lightllm/models/deepseek_mtp/model.py @@ -1,6 +1,6 @@ from typing import List from lightllm.models.deepseek2.model import Deepseek2TpPartModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.deepseek_mtp.layer_infer.pre_layer_infer import Deepseek3MTPPreLayerInfer from lightllm.models.deepseek_mtp.layer_weights.pre_and_post_layer_weight import Deepseek3MTPPreAndPostLayerWeight from lightllm.common.basemodel import TpPartBaseModel diff --git a/lightllm/models/draft_registry.py b/lightllm/models/draft_registry.py new file mode 100644 index 0000000000..fe56e694e4 --- /dev/null +++ b/lightllm/models/draft_registry.py @@ -0,0 +1,46 @@ +"""Registry for mapping target model types and speculative modes to draft models.""" + +from typing import Callable, Dict, List, Tuple, Type, TypeVar, Union + + +T = TypeVar("T") + + +class _DraftModelRegistry: + def __init__(self): + self._registry: Dict[Tuple[str, str], Type] = {} + + def __call__( + self, + model_type: Union[str, List[str], Tuple[str, ...]], + spec_modes: Union[str, List[str], Tuple[str, ...]], + ) -> Callable[[T], T]: + model_types = (model_type,) if isinstance(model_type, str) else tuple(model_type) + modes = (spec_modes,) if isinstance(spec_modes, str) else tuple(spec_modes) + + def decorator(model_class: T) -> T: + for current_model_type in model_types: + for spec_mode in modes: + key = (current_model_type, spec_mode) + if key in self._registry: + raise ValueError(f"Duplicate draft model registration: {key}") + self._registry[key] = model_class + return model_class + + return decorator + + def get_model_class(self, model_cfg: dict, spec_mode: str) -> Type: + model_type = model_cfg.get("model_type", "") + try: + return self._registry[(model_type, spec_mode)] + except KeyError: + raise ValueError( + f"Unsupported speculative draft model: mode={spec_mode}, model_type={model_type}" + ) from None + + +DraftModelRegistry = _DraftModelRegistry() + + +def get_draft_model_class(model_cfg: dict, spec_mode: str) -> Type: + return DraftModelRegistry.get_model_class(model_cfg=model_cfg, spec_mode=spec_mode) diff --git a/lightllm/models/glm4_moe_lite_mtp/model.py b/lightllm/models/glm4_moe_lite_mtp/model.py index 3218378026..06f941d806 100644 --- a/lightllm/models/glm4_moe_lite_mtp/model.py +++ b/lightllm/models/glm4_moe_lite_mtp/model.py @@ -1,7 +1,7 @@ from typing import List from lightllm.models.deepseek_mtp.layer_infer.pre_layer_infer import Deepseek3MTPPreLayerInfer from lightllm.models.glm4_moe_lite.model import Glm4MoeLiteTpPartModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.glm4_moe_lite_mtp.layer_weights.pre_and_post_layer_weight import ( Glm4MoeLiteMTPPreAndPostLayerWeight, ) diff --git a/lightllm/models/mistral_mtp/model.py b/lightllm/models/mistral_mtp/model.py index ae1e208382..6a3c32ebbc 100644 --- a/lightllm/models/mistral_mtp/model.py +++ b/lightllm/models/mistral_mtp/model.py @@ -1,6 +1,6 @@ from typing import List from lightllm.models.mistral.model import MistralTpPartModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.mistral_mtp.layer_weights.pre_and_post_layer_weight import MistralMTPPreAndPostLayerWeight from lightllm.models.mistral_mtp.layer_infer.pre_layer_infer import MistralMTPPreLayerInfer from lightllm.models.mistral_mtp.layer_infer.post_layer_infer import MistralMTPPostLayerInfer diff --git a/lightllm/models/qwen3_5_dflash/model.py b/lightllm/models/qwen3_5_dflash/model.py index da3755588d..ae376ea334 100644 --- a/lightllm/models/qwen3_5_dflash/model.py +++ b/lightllm/models/qwen3_5_dflash/model.py @@ -1,6 +1,6 @@ from lightllm.models.llama.model import LlamaTpPartModel from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_5_dflash.layer_weights.pre_and_post_layer_weight import ( Qwen35DFlashPreAndPostLayerWeight, ) diff --git a/lightllm/models/qwen3_5_dspark/model.py b/lightllm/models/qwen3_5_dspark/model.py index 23182ce0f6..aba188d310 100644 --- a/lightllm/models/qwen3_5_dspark/model.py +++ b/lightllm/models/qwen3_5_dspark/model.py @@ -1,6 +1,6 @@ from lightllm.models.llama.model import LlamaTpPartModel from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry @DraftModelRegistry(model_type=("qwen3_5", "qwen3_5_text"), spec_modes="dspark") diff --git a/lightllm/models/qwen3_5_moe_mtp/model.py b/lightllm/models/qwen3_5_moe_mtp/model.py index 9e45501597..e852e24c59 100644 --- a/lightllm/models/qwen3_5_moe_mtp/model.py +++ b/lightllm/models/qwen3_5_moe_mtp/model.py @@ -1,5 +1,5 @@ from lightllm.models.qwen3_5_mtp.model import Qwen3_5MTPModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_5_moe_mtp.layer_weights.transformer_layer_weight import ( Qwen3_5MoeMTPTransformerLayerWeight, ) diff --git a/lightllm/models/qwen3_5_mtp/model.py b/lightllm/models/qwen3_5_mtp/model.py index 760e08dd97..b8639a9970 100644 --- a/lightllm/models/qwen3_5_mtp/model.py +++ b/lightllm/models/qwen3_5_mtp/model.py @@ -2,7 +2,7 @@ from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.models.qwen3_5.model import Qwen3_5TpPartModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_5.layer_infer.transformer_layer_infer import Qwen35TransformerLayerInfer from lightllm.models.qwen3_5_mtp.layer_weights.pre_and_post_layer_weight import Qwen3_5MTPPreAndPostLayerWeight from lightllm.models.qwen3_5_mtp.layer_weights.transformer_layer_weight import Qwen3_5MTPTransformerLayerWeight diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py index e0df32bfd3..48106ac730 100644 --- a/lightllm/models/qwen3_dflash/model.py +++ b/lightllm/models/qwen3_dflash/model.py @@ -4,7 +4,7 @@ ) from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.models.llama.model import LlamaTpPartModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer from lightllm.models.qwen3_dflash.layer_infer.pre_layer_infer import Qwen3DFlashPreLayerInfer diff --git a/lightllm/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py index 115f817876..9e42d02d41 100644 --- a/lightllm/models/qwen3_dspark/model.py +++ b/lightllm/models/qwen3_dspark/model.py @@ -1,5 +1,5 @@ from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_dspark.infer_struct import Qwen3DSparkInferStateInfo from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput diff --git a/lightllm/models/qwen3_eagle/model.py b/lightllm/models/qwen3_eagle/model.py index 02706f0b28..83c90ebbde 100644 --- a/lightllm/models/qwen3_eagle/model.py +++ b/lightllm/models/qwen3_eagle/model.py @@ -3,7 +3,7 @@ from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.models.llama.model import LlamaTpPartModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_eagle.layer_infer.pre_layer_infer import Qwen3EaglePreLayerInfer from lightllm.models.qwen3_eagle.layer_infer.transformer_layer_infer import Qwen3EagleTransformerLayerInfer from lightllm.models.qwen3_eagle.layer_weights.pre_and_post_layer_weight import Qwen3EaglePreAndPostLayerWeight diff --git a/lightllm/models/qwen3_moe_mtp/model.py b/lightllm/models/qwen3_moe_mtp/model.py index a4b95a21fc..88522a5e23 100644 --- a/lightllm/models/qwen3_moe_mtp/model.py +++ b/lightllm/models/qwen3_moe_mtp/model.py @@ -1,6 +1,6 @@ from typing import List from lightllm.models.qwen3_moe.model import Qwen3MOEModel -from lightllm.models.registry import DraftModelRegistry +from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_moe_mtp.layer_weights.pre_and_post_layer_weight import Qwen3MOEMTPPreAndPostLayerWeight from lightllm.models.deepseek_mtp.layer_infer.pre_layer_infer import Deepseek3MTPPreLayerInfer from lightllm.models.qwen3_moe_mtp.layer_infer.transformer_layer_infer import Qwen3MOEMTPTransformerLayerInfer diff --git a/lightllm/models/registry.py b/lightllm/models/registry.py index b2e7ffa4ac..1b513cc27d 100644 --- a/lightllm/models/registry.py +++ b/lightllm/models/registry.py @@ -1,6 +1,6 @@ import collections from dataclasses import dataclass -from typing import Callable, Dict, List, Optional, Tuple, Type, TypeVar, Union +from typing import Callable, Dict, List, Optional, Type, TypeVar, Union from lightllm.utils.log_utils import init_logger @@ -89,42 +89,6 @@ def get_model_class(self, model_cfg: dict): ModelRegistry = _ModelRegistries() -class _DraftModelRegistry: - def __init__(self): - self._registry: Dict[Tuple[str, str], Type] = {} - - def __call__( - self, - model_type: Union[str, List[str], Tuple[str, ...]], - spec_modes: Union[str, List[str], Tuple[str, ...]], - ) -> Callable[[T], T]: - model_types = (model_type,) if isinstance(model_type, str) else tuple(model_type) - modes = (spec_modes,) if isinstance(spec_modes, str) else tuple(spec_modes) - - def decorator(model_class: T) -> T: - for current_model_type in model_types: - for spec_mode in modes: - key = (current_model_type, spec_mode) - if key in self._registry: - raise ValueError(f"Duplicate draft model registration: {key}") - self._registry[key] = model_class - return model_class - - return decorator - - def get_model_class(self, model_cfg: dict, spec_mode: str) -> Type: - model_type = model_cfg.get("model_type", "") - try: - return self._registry[(model_type, spec_mode)] - except KeyError: - raise ValueError( - f"Unsupported speculative draft model: mode={spec_mode}, model_type={model_type}" - ) from None - - -DraftModelRegistry = _DraftModelRegistry() - - def get_model(model_cfg: dict, model_kvargs: dict): try: model, is_multimodal = ModelRegistry.get_model(model_cfg, model_kvargs) @@ -143,10 +107,6 @@ def get_model_class(model_cfg: dict): raise -def get_draft_model_class(model_cfg, spec_mode): - return DraftModelRegistry.get_model_class(model_cfg=model_cfg, spec_mode=spec_mode) - - def is_reward_model() -> Callable[[Dict[str, any]], bool]: """Predicate: whether the model is RewardModel.""" return lambda model_cfg: "RewardModel" in model_cfg.get("architectures", [""])[0] From 9704c2044a334fa6cfa5ef6a6a948272f737d6a5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 00:59:53 +0000 Subject: [PATCH 028/103] refactor: rename dynamic MTP flag --- .../router/model_infer/mode_backend/base_backend.py | 2 +- .../model_infer/mode_backend/chunked_prefill/impl.py | 8 ++++---- .../server/router/model_infer/speculative/engine.py | 12 ++++++------ .../model_infer/speculative/proposers/__init__.py | 12 ++++++------ .../router/model_infer/speculative/proposers/base.py | 4 ++-- .../model_infer/speculative/proposers/dflash.py | 4 ++-- .../model_infer/speculative/proposers/dspark.py | 4 ++-- .../model_infer/speculative/proposers/eagle_mtp.py | 2 +- .../model_infer/speculative/proposers/vanilla_mtp.py | 4 ++-- .../model_infer/speculative/test_eagle_overlap.py | 4 ++-- .../router/model_infer/speculative/test_planner.py | 12 ++++++------ unit_tests/utils/test_speculative_utils.py | 6 +++--- 12 files changed, 37 insertions(+), 37 deletions(-) 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 e8308af36b..c0a88a7e38 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -303,7 +303,7 @@ def init_spec_engine(self, main_kvargs: dict): self.spec_engine = SpecEngine( backend=self, spec_mode=self.args.mtp_mode, - enable_dynamic_spec=self.args.mtp_dynamic_verify, + enable_dynmaic_mtp=self.args.mtp_dynamic_verify, ) return diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 22bcee68fc..75b6b5357b 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -28,13 +28,13 @@ def __init__(self) -> None: # 用于控制每一步是执行prefill 和 decode 还是跳过 self.control_state_machine = ControlState() - self.enable_dynamic_spec = False + self.enable_dynmaic_mtp = False # 在 mtp 模式下切换绑定的prefill 和 decode 函数 if get_env_start_args().mtp_mode is not None: self.prefill = self.prefill_mtp self.decode = self.decode_mtp - self.enable_dynamic_spec = get_env_start_args().mtp_dynamic_verify + self.enable_dynmaic_mtp = get_env_start_args().mtp_dynamic_verify else: self.prefill = self.prefill_normal self.decode = self.decode_normal @@ -44,7 +44,7 @@ def __init__(self) -> None: def init_custom(self): super().init_custom() - if self.enable_dynamic_spec: + if self.enable_dynmaic_mtp: self.spec_gloo_group = create_new_group_for_current_dp("gloo") logger.info(f"spec_gloo_group ranks {dist.get_rank(self.spec_gloo_group)}") @@ -356,7 +356,7 @@ def decode_mtp( spec_engine.record_request_spec_metrics( decode_reqs=decode_reqs, accept_lengths_cpu=mtp_accept_len_cpu, - verified_row_reqs=run_reqs if self.enable_dynamic_spec else None, + verified_row_reqs=run_reqs if self.enable_dynmaic_mtp else None, ) select_mask = accepted_index_cpu.to(dtype=torch.bool) diff --git a/lightllm/server/router/model_infer/speculative/engine.py b/lightllm/server/router/model_infer/speculative/engine.py index e68b01fa34..0e3423ec54 100644 --- a/lightllm/server/router/model_infer/speculative/engine.py +++ b/lightllm/server/router/model_infer/speculative/engine.py @@ -30,14 +30,14 @@ class SpecEngine: state and proposal generation. """ - def __init__(self, backend, spec_mode: str, enable_dynamic_spec: bool) -> None: + def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: self.backend = backend self.spec_mode = spec_mode - self.enable_dynamic_spec = enable_dynamic_spec + self.enable_dynmaic_mtp = enable_dynmaic_mtp self.proposer = build_spec_proposer( spec_mode=spec_mode, backend=backend, - enable_dynamic_spec=enable_dynamic_spec, + enable_dynmaic_mtp=enable_dynmaic_mtp, ) self.planner = self._build_decode_planner() self._register_cuda_graph_costs() @@ -249,7 +249,7 @@ def update_planner_feedback( ) -> None: """Feed iteration-level observations into the dynamic planner.""" - if not self.enable_dynamic_spec: + if not self.enable_dynmaic_mtp: return self.planner.update_feedback( @@ -306,7 +306,7 @@ def free_unused_decode_mem( # Construction helpers. def _build_decode_planner(self): - if not self.enable_dynamic_spec: + if not self.enable_dynmaic_mtp: return FixedSpecPlanner(max_draft_step=self.backend.max_draft_step) if self.spec_mode == "dspark": return DSparkPlanner( @@ -320,7 +320,7 @@ def _build_decode_planner(self): ) def _register_cuda_graph_costs(self) -> None: - if not self.enable_dynamic_spec: + if not self.enable_dynmaic_mtp: return target_graph = self.backend.model.graph diff --git a/lightllm/server/router/model_infer/speculative/proposers/__init__.py b/lightllm/server/router/model_infer/speculative/proposers/__init__.py index 10f3e31b47..f7a19bce61 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/__init__.py +++ b/lightllm/server/router/model_infer/speculative/proposers/__init__.py @@ -4,27 +4,27 @@ from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer -def build_spec_proposer(*, spec_mode: str, backend, enable_dynamic_spec: bool) -> "BaseSpecProposer": +def build_spec_proposer(*, spec_mode: str, backend, enable_dynmaic_mtp: bool) -> "BaseSpecProposer": if spec_mode == "dspark": from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer - return DSparkProposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) + return DSparkProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) if spec_mode == "dflash": from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer - return DFlashProposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) + return DFlashProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) if spec_mode == "eagle3": from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer - return Eagle3Proposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) + return Eagle3Proposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) if spec_mode in ("eagle_with_att", "eagle_no_att"): from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer - return EagleMTPProposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) + return EagleMTPProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) if spec_mode in ("vanilla_with_att", "vanilla_no_att"): from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer - return VanillaMTPProposer(backend=backend, enable_dynamic_spec=enable_dynamic_spec) + return VanillaMTPProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) raise ValueError(f"unsupported speculative mode: {spec_mode}") diff --git a/lightllm/server/router/model_infer/speculative/proposers/base.py b/lightllm/server/router/model_infer/speculative/proposers/base.py index b61847ee01..6f9d9e9145 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/speculative/proposers/base.py @@ -41,9 +41,9 @@ class BaseSpecProposer: but does not verify acceptance; verification is handled by SpecEngine. """ - def __init__(self, *, backend: "ModeBackend", enable_dynamic_spec: bool) -> None: + def __init__(self, *, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> None: self.backend = backend - self.enable_dynamic_spec = bool(enable_dynamic_spec) + self.enable_dynmaic_mtp = bool(enable_dynmaic_mtp) def get_draft_steps(self) -> Tuple[int, ...]: """Return the draft configurations supported by this proposer.""" diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/speculative/proposers/dflash.py index c700033e21..5551de26b9 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dflash.py @@ -48,7 +48,7 @@ def propose_next( ) draft_output = draft_model.forward(draft_input) - if self.enable_dynamic_spec: + if self.enable_dynmaic_mtp: flat_draft_token_ids, flat_draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) else: flat_draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) @@ -56,7 +56,7 @@ def propose_next( proposal_token_ids[accepted_tail_rows, 1:] = block_draft_token_ids[:, :draft_step] schedule_scores = None - if self.enable_dynamic_spec: + if self.enable_dynmaic_mtp: block_draft_token_probs = flat_draft_token_probs.reshape(request_count, block_size) schedule_scores = torch.zeros( (verify_row_count, draft_step), diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/speculative/proposers/dspark.py index 641cade2d0..6e17a02cc7 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/speculative/proposers/dspark.py @@ -41,7 +41,7 @@ def propose_next( dtype=torch.float32, device=next_token_ids.device, ) - if self.enable_dynamic_spec + if self.enable_dynmaic_mtp else None ) @@ -65,7 +65,7 @@ def propose_next( block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) proposal_token_ids[accepted_tail_rows, 1:] = block_draft_token_ids[:, :draft_step] - if self.enable_dynamic_spec: + if self.enable_dynmaic_mtp: confidence_logits = draft_output.confidence_logits if confidence_logits is None: raise RuntimeError("DSpark dynamic verify requires confidence head logits") diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py index 8814e07e99..01fcedbd3d 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py @@ -127,7 +127,7 @@ def propose_next( fill_value=1, ) proposal_token_ids[:, 0].copy_(next_token_ids) - collect_schedule_scores = self.enable_dynamic_spec + collect_schedule_scores = self.enable_dynmaic_mtp schedule_scores = ( torch.zeros( (verify_row_count, draft_step), diff --git a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py index c044e53874..43a99c3254 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py @@ -99,7 +99,7 @@ def propose_next( dtype=torch.float32, device=next_token_ids.device, ) - if self.enable_dynamic_spec + if self.enable_dynmaic_mtp else None ) @@ -109,7 +109,7 @@ def propose_next( main_model_input.mtp_draft_input_hiddens = draft_hidden draft_output = draft_model.forward(main_model_input) draft_hidden = draft_output.spec_hidden - if self.enable_dynamic_spec: + if self.enable_dynmaic_mtp: draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) schedule_scores[:, step] = draft_token_probs else: diff --git a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py index d927fb2678..bfa3bc1ea6 100644 --- a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py @@ -71,7 +71,7 @@ def test_overlap_eagle_keeps_fixed_verify_layout(): ), _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) - proposer = EagleMTPProposer(backend=backend, enable_dynamic_spec=False) + proposer = EagleMTPProposer(backend=backend, enable_dynmaic_mtp=False) proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) model_input0 = _target_input(batch_size=6) model_input1 = _target_input(batch_size=6) @@ -114,7 +114,7 @@ def test_autoregressive_eagle_reuses_overlap_inputs(): ), _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) - proposer = AutoregressiveEagleProposer(backend=backend, enable_dynamic_spec=False) + proposer = AutoregressiveEagleProposer(backend=backend, enable_dynmaic_mtp=False) proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) model_input0 = _target_input(batch_size=6) model_input1 = _target_input(batch_size=6) diff --git a/unit_tests/server/router/model_infer/speculative/test_planner.py b/unit_tests/server/router/model_infer/speculative/test_planner.py index 920ada9a14..27fca26aaa 100644 --- a/unit_tests/server/router/model_infer/speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/speculative/test_planner.py @@ -28,7 +28,7 @@ def build_draft_cost_provider(proposer_class, max_draft_step: int = 3, block_siz max_draft_step=max_draft_step, draft_models=[SimpleNamespace(block_size=block_size)], ) - return proposer_class(backend=backend, enable_dynamic_spec=True) + return proposer_class(backend=backend, enable_dynmaic_mtp=True) def build_lightspec_planner( @@ -57,10 +57,10 @@ def build_dspark_planner(max_draft_step: int = 3, block_size: int = 3): ) -def build_planner(spec_mode: str, enable_dynamic_spec: bool = True): +def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): engine = SpecEngine.__new__(SpecEngine) engine.spec_mode = spec_mode - engine.enable_dynamic_spec = enable_dynamic_spec + engine.enable_dynmaic_mtp = enable_dynmaic_mtp engine.backend = SimpleNamespace( max_draft_step=3, draft_models=[SimpleNamespace(block_size=3)], @@ -92,7 +92,7 @@ def test_infer_cost_candidates_include_feasible_boundaries(): def test_engine_routes_only_dspark_to_the_confidence_planner(): - assert isinstance(build_planner("eagle3", enable_dynamic_spec=False), FixedSpecPlanner) + assert isinstance(build_planner("eagle3", enable_dynmaic_mtp=False), FixedSpecPlanner) assert isinstance(build_planner("dspark"), DSparkPlanner) dflash_planner = build_planner("dflash") @@ -453,7 +453,7 @@ def test_lightspec_short_current_proposal_can_recover_to_a_deeper_draft(): def test_engine_skips_feedback_for_a_mixed_proposal_batch(): engine = SpecEngine.__new__(SpecEngine) - engine.enable_dynamic_spec = True + engine.enable_dynmaic_mtp = True engine.planner = build_lightspec_planner() plan = SpecDecodePlan( dynamic_batch_size=5, @@ -488,7 +488,7 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): schedule_scores_cpu=torch.from_numpy(confidence_probs), ) engine = SpecEngine.__new__(SpecEngine) - engine.enable_dynamic_spec = True + engine.enable_dynmaic_mtp = True engine.planner = planner engine.update_planner_feedback( diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index f2fa663e7e..83015b81bf 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -236,7 +236,7 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): draft_models=[draft_model], _gen_argmax_token_ids_and_prob=lambda _: (flat_draft_token_ids, flat_draft_token_probs), ) - proposer = DFlashProposer(backend=backend, enable_dynamic_spec=True) + proposer = DFlashProposer(backend=backend, enable_dynmaic_mtp=True) proposer.extend_draft_kv_cache = lambda **_: None proposer.build_block_draft_input = lambda **_: (SimpleNamespace(), torch.tensor([10, 11])) @@ -264,7 +264,7 @@ def test_dflash_reuses_decode_input_for_kv_commit(): draft_model = SimpleNamespace(forward=forwarded_inputs.append) proposer = DFlashProposer( backend=SimpleNamespace(draft_models=[draft_model]), - enable_dynamic_spec=False, + enable_dynmaic_mtp=False, ) input_ids = torch.arange(4) multimodal_params = [{"images": [], "audios": []}] * 4 @@ -302,7 +302,7 @@ def test_dflash_expands_position_delta_with_request_block_rows(): backend=SimpleNamespace( draft_models=[SimpleNamespace(block_size=block_size, mask_token_id=99)], ), - enable_dynamic_spec=False, + enable_dynmaic_mtp=False, ) proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) model_input = SimpleNamespace( From 7b708d581024a1dfe9d68235f8b86ae93399b525 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 01:04:46 +0000 Subject: [PATCH 029/103] refactor: remove unused MTP Gloo group --- .../model_infer/mode_backend/chunked_prefill/impl.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 75b6b5357b..8661208b18 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -1,6 +1,5 @@ import torch import time -import torch.distributed as dist from typing import List from lightllm.common.basemodel.triton_kernel.mtp_utils import gen_b_req_mtp_start_loc from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend @@ -16,7 +15,6 @@ from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_current_device_id from .control_state import ControlState -from lightllm.utils.dist_utils import create_new_group_for_current_dp from lightllm.utils.envs_utils import get_env_start_args logger = init_logger(__name__) @@ -42,12 +40,6 @@ def __init__(self) -> None: self.classed_req_strict_prefill = False return - def init_custom(self): - super().init_custom() - if self.enable_dynmaic_mtp: - self.spec_gloo_group = create_new_group_for_current_dp("gloo") - logger.info(f"spec_gloo_group ranks {dist.get_rank(self.spec_gloo_group)}") - def infer_loop(self): torch.cuda.set_device(get_current_device_id()) try: From d3e657fc544f0def6933eb2d59dd4a7982f40fb5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 01:20:57 +0000 Subject: [PATCH 030/103] refactor: build MTP group markers from request IDs --- .../generic_padded_pre_process.py | 18 +++++++- .../mode_backend/generic_pre_process.py | 46 ++++++++++++------- .../mode_backend/test_generic_pre_process.py | 14 +++--- 3 files changed, 52 insertions(+), 26 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py index 757288a508..0f13934e57 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py @@ -15,7 +15,7 @@ from .generic_pre_process import ( build_b_position_delta, build_diverse_shared_group_infos, - build_spec_shared_group_markers, + build_mtp_shared_group_markers, ) @@ -215,7 +215,21 @@ def padded_prepare_decode_inputs( b_mark_shared_group = F.pad(b_mark_shared_group, (0, padded_row_count), value=1) elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): b_shared_seq_len = None - b_mark_shared_group = build_spec_shared_group_markers(b_mtp_index=b_mtp_index) + real_row_count = b_req_idx.shape[0] - padded_row_count + b_mark_shared_group = build_mtp_shared_group_markers(b_req_idx=b_req_idx[:real_row_count]) + if padded_row_count > 0: + # ModelInput must use HOLD_REQUEST_ID for every padded row, which + # makes adjacent fake requests indistinguishable by b_req_idx. Use + # local synthetic request ids only to preserve their MTP group + # boundaries while building the marker tensor. + mtp_size = args_mtp_step + 1 + padded_b_req_idx = torch.arange( + padded_req_num, + dtype=b_req_idx.dtype, + device=b_req_idx.device, + ).repeat_interleave(mtp_size) + padded_group_markers = build_mtp_shared_group_markers(b_req_idx=padded_b_req_idx) + b_mark_shared_group = torch.cat((b_mark_shared_group, padded_group_markers)) else: b_shared_seq_len = None b_mark_shared_group = None diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 3bfdc8b063..21dbb02cff 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -138,7 +138,7 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): b_shared_seq_len = None - b_mark_shared_group = build_spec_shared_group_markers(b_mtp_index=b_mtp_index) + b_mark_shared_group = build_mtp_shared_group_markers(b_req_idx=b_req_idx) else: b_shared_seq_len = None b_mark_shared_group = None @@ -224,24 +224,38 @@ def build_diverse_shared_group_infos(run_reqs: List[InferReq]) -> Tuple[torch.Te return b_shared_seq_len, b_mark_shared_group -def build_spec_shared_group_markers(b_mtp_index: torch.Tensor) -> torch.Tensor: - # Each logical request starts at row index 0. Only the final row of each - # speculative query group stores its group size; earlier rows store zero. +def build_mtp_shared_group_markers(b_req_idx: torch.Tensor) -> torch.Tensor: + """Build MTP row-group markers from consecutive request indexes. + + Rows belonging to the same request form one MTP group. Only the final row + of a group stores the group size; all earlier rows store zero. A request + with more rows than the attention kernel supports is split into multiple + adjacent groups. + + Example with ``max_batch_shared_group_size == 3``:: + + b_req_idx: [7, 7, 7, 7, 11, 11, 20] + MTP groups: [7, 7, 7] [7] [11, 11] [20] + b_mark_shared_group: [0, 0, 3, 1, 0, 2, 1] + + The request id itself is not written to the result; it is only used to + detect where one request ends and the next request begins. + """ max_batch_shared_group_size = get_diverse_max_batch_shared_group_size() - mtp_indexes = b_mtp_index.tolist() - b_mark_shared_group = [] + assert max_batch_shared_group_size > 0 + + req_indexes = b_req_idx.tolist() + b_mark_shared_group = [0] * len(req_indexes) group_start = 0 - for group_end in range(1, len(mtp_indexes) + 1): - reaches_request_boundary = group_end == len(mtp_indexes) or mtp_indexes[group_end] == 0 - reaches_size_limit = group_end - group_start == max_batch_shared_group_size - if not reaches_request_boundary and not reaches_size_limit: + for row_index, req_idx in enumerate(req_indexes): + is_request_end = row_index == len(req_indexes) - 1 or req_indexes[row_index + 1] != req_idx + group_size = row_index - group_start + 1 + reaches_size_limit = group_size == max_batch_shared_group_size + if not is_request_end and not reaches_size_limit: continue - group_size = group_end - group_start - b_mark_shared_group.extend([0] * (group_size - 1)) - b_mark_shared_group.append(group_size) - group_start = group_end + b_mark_shared_group[row_index] = group_size + group_start = row_index + 1 - assert len(b_mark_shared_group) == len(mtp_indexes) - b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cpu") + b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device=b_req_idx.device) return b_mark_shared_group diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index ad3d9bccfe..195c7c1d3e 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -4,22 +4,20 @@ from lightllm.server.router.model_infer.mode_backend import generic_padded_pre_process, generic_pre_process -def test_spec_shared_group_markers_split_requests_and_size_limit(monkeypatch): +def test_mtp_shared_group_markers_split_requests_and_size_limit(monkeypatch): monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 2) - b_mtp_index = torch.tensor([0, 1, 2, 0, 1, 0], dtype=torch.int32) + b_req_idx = torch.tensor([7, 7, 7, 11, 11, 20], dtype=torch.int32) - markers = generic_pre_process.build_spec_shared_group_markers(b_mtp_index=b_mtp_index) + markers = generic_pre_process.build_mtp_shared_group_markers(b_req_idx=b_req_idx) assert markers.tolist() == [0, 2, 1, 0, 2, 1] -def test_spec_shared_group_markers_include_padded_request_rows(monkeypatch): +def test_mtp_shared_group_markers_detect_request_boundaries(monkeypatch): monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 8) - # The last two groups represent padded requests. Their request ids are both - # HOLD_REQUEST_ID, so mtp_index resets define the actual group boundaries. - b_mtp_index = torch.tensor([0, 1, 2, 0, 1, 2, 0, 1, 2], dtype=torch.int32) + b_req_idx = torch.tensor([7, 7, 7, 11, 11, 11, 20, 20, 20], dtype=torch.int32) - markers = generic_pre_process.build_spec_shared_group_markers(b_mtp_index=b_mtp_index) + markers = generic_pre_process.build_mtp_shared_group_markers(b_req_idx=b_req_idx) assert markers.tolist() == [0, 0, 3, 0, 0, 3, 0, 0, 3] From dca9feae9467589c6eeedc826d0e01fd28c4d47d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 01:30:06 +0000 Subject: [PATCH 031/103] refactor: simplify padded MTP group markers --- .../mode_backend/generic_padded_pre_process.py | 16 +--------------- .../mode_backend/test_generic_pre_process.py | 2 +- 2 files changed, 2 insertions(+), 16 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py index 0f13934e57..1ebd95e3d3 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py @@ -215,21 +215,7 @@ def padded_prepare_decode_inputs( b_mark_shared_group = F.pad(b_mark_shared_group, (0, padded_row_count), value=1) elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): b_shared_seq_len = None - real_row_count = b_req_idx.shape[0] - padded_row_count - b_mark_shared_group = build_mtp_shared_group_markers(b_req_idx=b_req_idx[:real_row_count]) - if padded_row_count > 0: - # ModelInput must use HOLD_REQUEST_ID for every padded row, which - # makes adjacent fake requests indistinguishable by b_req_idx. Use - # local synthetic request ids only to preserve their MTP group - # boundaries while building the marker tensor. - mtp_size = args_mtp_step + 1 - padded_b_req_idx = torch.arange( - padded_req_num, - dtype=b_req_idx.dtype, - device=b_req_idx.device, - ).repeat_interleave(mtp_size) - padded_group_markers = build_mtp_shared_group_markers(b_req_idx=padded_b_req_idx) - b_mark_shared_group = torch.cat((b_mark_shared_group, padded_group_markers)) + b_mark_shared_group = build_mtp_shared_group_markers(b_req_idx=b_req_idx) else: b_shared_seq_len = None b_mark_shared_group = None diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index 195c7c1d3e..8c8ff37d92 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -56,4 +56,4 @@ def test_padded_decode_builds_spec_metadata_for_real_and_fake_rows(monkeypatch): assert padded_req_num == 2 assert model_input.b_mtp_index.tolist() == [0, 1, 2, 0, 1, 2, 0, 1, 2] - assert model_input.b_mark_shared_group.tolist() == [0, 0, 3, 0, 0, 3, 0, 0, 3] + assert model_input.b_mark_shared_group.tolist() == [0, 0, 3, 0, 0, 0, 0, 0, 6] From c82f2837f808d25e6f39d33801f5c870fb31d848 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 01:33:16 +0000 Subject: [PATCH 032/103] refactor: rename speculative package --- .../model_infer/mode_backend/base_backend.py | 2 +- .../model_infer/mtp_speculative/__init__.py | 4 ++++ .../{speculative => mtp_speculative}/engine.py | 6 +++--- .../planner.py | 2 +- .../proposers/__init__.py | 12 ++++++------ .../proposers/base.py | 0 .../proposers/dflash.py | 4 ++-- .../proposers/dspark.py | 4 ++-- .../proposers/eagle3.py | 2 +- .../proposers/eagle_mtp.py | 2 +- .../proposers/parallel_block.py | 2 +- .../proposers/vanilla_mtp.py | 2 +- .../router/model_infer/speculative/__init__.py | 4 ---- .../mode_backend/test_dp_spec_engine.py | 2 +- .../test_eagle_overlap.py | 4 ++-- .../test_planner.py | 18 +++++++++--------- unit_tests/utils/test_speculative_utils.py | 2 +- 17 files changed, 36 insertions(+), 36 deletions(-) create mode 100644 lightllm/server/router/model_infer/mtp_speculative/__init__.py rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/engine.py (98%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/planner.py (99%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/proposers/__init__.py (59%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/proposers/base.py (100%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/proposers/dflash.py (93%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/proposers/dspark.py (94%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/proposers/eagle3.py (79%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/proposers/eagle_mtp.py (99%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/proposers/parallel_block.py (98%) rename lightllm/server/router/model_infer/{speculative => mtp_speculative}/proposers/vanilla_mtp.py (97%) delete mode 100644 lightllm/server/router/model_infer/speculative/__init__.py rename unit_tests/server/router/model_infer/{speculative => mtp_speculative}/test_eagle_overlap.py (97%) rename unit_tests/server/router/model_infer/{speculative => mtp_speculative}/test_planner.py (95%) 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 c0a88a7e38..1a0e867d4e 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -36,7 +36,7 @@ get_radix_tree_merge_update_delta, ) from lightllm.distributed import dist_group_manager -from lightllm.server.router.model_infer.speculative import SpecEngine +from lightllm.server.router.model_infer.mtp_speculative import SpecEngine from lightllm.distributed.communication_op import ( all_gather_into_tensor, all_reduce, diff --git a/lightllm/server/router/model_infer/mtp_speculative/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/__init__.py new file mode 100644 index 0000000000..779438e356 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/__init__.py @@ -0,0 +1,4 @@ +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine + + +__all__ = ["SpecEngine"] diff --git a/lightllm/server/router/model_infer/speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py similarity index 98% rename from lightllm/server/router/model_infer/speculative/engine.py rename to lightllm/server/router/model_infer/mtp_speculative/engine.py index 0e3423ec54..73de0e73e1 100644 --- a/lightllm/server/router/model_infer/speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -12,14 +12,14 @@ mtp_verify, ) from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager -from lightllm.server.router.model_infer.speculative.planner import ( +from lightllm.server.router.model_infer.mtp_speculative.planner import ( DSparkPlanner, FixedSpecPlanner, LightSpecPlanner, SpecDecodePlan, ) -from lightllm.server.router.model_infer.speculative.proposers import build_spec_proposer -from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers import build_spec_proposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal class SpecEngine: diff --git a/lightllm/server/router/model_infer/speculative/planner.py b/lightllm/server/router/model_infer/mtp_speculative/planner.py similarity index 99% rename from lightllm/server/router/model_infer/speculative/planner.py rename to lightllm/server/router/model_infer/mtp_speculative/planner.py index 5b24701c83..41e18106da 100644 --- a/lightllm/server/router/model_infer/speculative/planner.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner.py @@ -8,7 +8,7 @@ from sortedcontainers import SortedDict if TYPE_CHECKING: - from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer @dataclass(frozen=True) diff --git a/lightllm/server/router/model_infer/speculative/proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py similarity index 59% rename from lightllm/server/router/model_infer/speculative/proposers/__init__.py rename to lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py index f7a19bce61..e4db7a9732 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py @@ -1,28 +1,28 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: - from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer def build_spec_proposer(*, spec_mode: str, backend, enable_dynmaic_mtp: bool) -> "BaseSpecProposer": if spec_mode == "dspark": - from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer + from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer return DSparkProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) if spec_mode == "dflash": - from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer + from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer return DFlashProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) if spec_mode == "eagle3": - from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer + from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer return Eagle3Proposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) if spec_mode in ("eagle_with_att", "eagle_no_att"): - from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import EagleMTPProposer + from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import EagleMTPProposer return EagleMTPProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) if spec_mode in ("vanilla_with_att", "vanilla_no_att"): - from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer + from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_mtp import VanillaMTPProposer return VanillaMTPProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) diff --git a/lightllm/server/router/model_infer/speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py similarity index 100% rename from lightllm/server/router/model_infer/speculative/proposers/base.py rename to lightllm/server/router/model_infer/mtp_speculative/proposers/base.py diff --git a/lightllm/server/router/model_infer/speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py similarity index 93% rename from lightllm/server/router/model_infer/speculative/proposers/dflash.py rename to lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index 5551de26b9..07dbcfbfe9 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -3,8 +3,8 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.speculative.proposers.parallel_block import ParallelBlockProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block import ParallelBlockProposer class DFlashProposer(ParallelBlockProposer): diff --git a/lightllm/server/router/model_infer/speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py similarity index 94% rename from lightllm/server/router/model_infer/speculative/proposers/dspark.py rename to lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 6e17a02cc7..2e981fc7a3 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -4,8 +4,8 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.speculative.proposers.parallel_block import ParallelBlockProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block import ParallelBlockProposer class DSparkProposer(ParallelBlockProposer): diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py similarity index 79% rename from lightllm/server/router/model_infer/speculative/proposers/eagle3.py rename to lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py index 453eaa5738..257a714118 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py @@ -2,7 +2,7 @@ import torch -from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import AutoregressiveEagleProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import AutoregressiveEagleProposer class Eagle3Proposer(AutoregressiveEagleProposer): diff --git a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py similarity index 99% rename from lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py rename to lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py index 01fcedbd3d..5245ec2cd1 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py @@ -5,7 +5,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal class AutoregressiveEagleProposer(BaseSpecProposer): diff --git a/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py similarity index 98% rename from lightllm/server/router/model_infer/speculative/proposers/parallel_block.py rename to lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py index 73c893b4e7..1b345f01cc 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/parallel_block.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py @@ -5,7 +5,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer class ParallelBlockProposer(BaseSpecProposer): diff --git a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py similarity index 97% rename from lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py rename to lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py index 43a99c3254..a8176392c8 100644 --- a/lightllm/server/router/model_infer/speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py @@ -3,7 +3,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal class VanillaMTPProposer(BaseSpecProposer): diff --git a/lightllm/server/router/model_infer/speculative/__init__.py b/lightllm/server/router/model_infer/speculative/__init__.py deleted file mode 100644 index 0697431885..0000000000 --- a/lightllm/server/router/model_infer/speculative/__init__.py +++ /dev/null @@ -1,4 +0,0 @@ -from lightllm.server.router.model_infer.speculative.engine import SpecEngine - - -__all__ = ["SpecEngine"] diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py index 0aa3f68031..1aef6b74b6 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py @@ -3,7 +3,7 @@ import torch from lightllm.server.router.model_infer.mode_backend.dp_backend.impl import DPChunkedPrefillBackend -from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal class _RecordingSpecEngine: diff --git a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py similarity index 97% rename from unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py rename to unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index bfa3bc1ea6..9ad007fb4b 100644 --- a/unit_tests/server/router/model_infer/speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -3,8 +3,8 @@ import torch from lightllm.common.basemodel.batch_objs import ModelOutput -from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import ( +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import ( AutoregressiveEagleProposer, EagleMTPProposer, ) diff --git a/unit_tests/server/router/model_infer/speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py similarity index 95% rename from unit_tests/server/router/model_infer/speculative/test_planner.py rename to unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 27fca26aaa..f6bb232d61 100644 --- a/unit_tests/server/router/model_infer/speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -3,24 +3,24 @@ import numpy as np import torch -from lightllm.server.router.model_infer.speculative.engine import SpecEngine -from lightllm.server.router.model_infer.speculative.planner import ( +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine +from lightllm.server.router.model_infer.mtp_speculative.planner import ( DSparkPlanner, FixedSpecPlanner, LightSpecPlanner, SpecDecodePlan, _InferCostMsTable, ) -from lightllm.server.router.model_infer.speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer -from lightllm.server.router.model_infer.speculative.proposers.dspark import DSparkProposer -from lightllm.server.router.model_infer.speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.speculative.proposers.eagle_mtp import ( +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import ( AutoregressiveEagleProposer, EagleMTPProposer, ) -from lightllm.server.router.model_infer.speculative.proposers.parallel_block import ParallelBlockProposer -from lightllm.server.router.model_infer.speculative.proposers.vanilla_mtp import VanillaMTPProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block import ParallelBlockProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_mtp import VanillaMTPProposer def build_draft_cost_provider(proposer_class, max_draft_step: int = 3, block_size: int = 3): diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 83015b81bf..3d7274187b 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -12,7 +12,7 @@ from lightllm.common.basemodel.attention.fa3.mla import MlaFa3DecodeAttState, MlaFa3PrefillAttState from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight -from lightllm.server.router.model_infer.speculative.proposers.dflash import DFlashProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer from lightllm.utils import envs_utils From ef43e5886c288f9cb7a3d56ef0981c7d6fafa9b6 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 01:36:26 +0000 Subject: [PATCH 033/103] refactor: align MTP metric method names --- lightllm/server/router/model_infer/infer_batch.py | 6 +++--- .../server/router/model_infer/mtp_speculative/engine.py | 6 +++--- .../router/model_infer/mtp_speculative/test_planner.py | 6 +++--- 3 files changed, 9 insertions(+), 9 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 976a20f65b..896e6afee6 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -855,14 +855,14 @@ def set_next_gen_token_id(self, next_token_id: int, logprob: float, output_len: self.shm_req.shm_logprobs.arr[index - 1] = (logprob, rank) return - def update_spec_accepted_token_num(self, accept_token_num: int): + def update_mtp_accepted_token_num(self, accept_token_num: int): # 用于统计 mtp 的接受率 self.shm_req.mtp_accepted_token_num += accept_token_num - def update_spec_verify_token_num(self, verify_token_num: int): + def update_mtp_verify_token_num(self, verify_token_num: int): self.shm_req.mtp_verify_token_num += verify_token_num - def update_spec_verify_step_num(self, verify_step_num: int): + def update_mtp_verify_step_num(self, verify_step_num: int): self.shm_req.mtp_verify_step_num += verify_step_num def get_last_gen_token(self): diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 73de0e73e1..b51e09b957 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -274,11 +274,11 @@ def record_request_spec_metrics( assert len(accept_lengths) == len(decode_reqs) verify_rows_by_req = None if verified_row_reqs is None else Counter(req.req_idx for req in verified_row_reqs) for req, accept_len in zip(decode_reqs, accept_lengths): - req.update_spec_accepted_token_num(accept_token_num=accept_len - 1) + req.update_mtp_accepted_token_num(accept_token_num=accept_len - 1) verify_token_num = req.mtp_step + 1 if verify_rows_by_req is None else verify_rows_by_req[req.req_idx] if verify_token_num > 0: - req.update_spec_verify_token_num(verify_token_num=verify_token_num) - req.update_spec_verify_step_num(verify_step_num=1) + req.update_mtp_verify_token_num(verify_token_num=verify_token_num) + req.update_mtp_verify_step_num(verify_step_num=1) def free_unused_decode_mem( self, diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index f6bb232d61..f440609e97 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -154,13 +154,13 @@ def __init__(self, req_idx: int, mtp_step: int): self.verified = 0 self.verify_steps = 0 - def update_spec_accepted_token_num(self, accept_token_num: int): + def update_mtp_accepted_token_num(self, accept_token_num: int): self.accepted += accept_token_num - def update_spec_verify_token_num(self, verify_token_num: int): + def update_mtp_verify_token_num(self, verify_token_num: int): self.verified += verify_token_num - def update_spec_verify_step_num(self, verify_step_num: int): + def update_mtp_verify_step_num(self, verify_step_num: int): self.verify_steps += verify_step_num engine = SpecEngine.__new__(SpecEngine) From e5f81a028aaee7ae72d36cab2f333ea2cb8f32e7 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 02:18:29 +0000 Subject: [PATCH 034/103] refactor: make MTP KV layer counts explicit --- lightllm/utils/envs_utils.py | 32 ++++++++++++++++------ unit_tests/utils/test_speculative_utils.py | 21 ++++++++++++++ 2 files changed, 45 insertions(+), 8 deletions(-) diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 917a7c9add..b3865b7af2 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -252,22 +252,38 @@ def enable_cpu_cache_numa_interleave() -> bool: @lru_cache(maxsize=None) def get_added_mtp_kv_layer_num() -> int: args = get_env_start_args() - if args.mtp_mode == "eagle_with_att": - return 1 - if args.mtp_mode == "vanilla_with_att": - return args.mtp_step - if args.mtp_mode not in ("eagle3", "dspark", "dflash"): + mtp_mode = args.mtp_mode + + if mtp_mode is None: + return 0 + if mtp_mode == "vanilla_no_att": + return 0 + if mtp_mode == "eagle_no_att": return 0 + if mtp_mode == "vanilla_with_att": + return args.mtp_step + if mtp_mode == "eagle_with_att": + return 1 + if mtp_mode == "eagle3": + return _get_mtp_draft_backbone_layer_num(args.mtp_draft_model_dir[0]) + if mtp_mode == "dspark": + return _get_mtp_draft_backbone_layer_num(args.mtp_draft_model_dir[0]) + if mtp_mode == "dflash": + return _get_mtp_draft_backbone_layer_num(args.mtp_draft_model_dir[0]) + + raise ValueError(f"unsupported mtp_mode: {mtp_mode}") + - draft_model_dir = args.mtp_draft_model_dir[0] +def _get_mtp_draft_backbone_layer_num(draft_model_dir: str) -> int: with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: draft_config = json.load(json_file) # Use the effective draft backbone config when the checkpoint stores it nested. draft_config.update(draft_config.get("dflash_config", {})) # A draft model may contain multiple attention layers; each layer needs a # separate KV-cache slot after the target model's layers. - total_added_mtp_kv_layer_num = draft_config.get("num_hidden_layers", draft_config.get("n_layer")) - return int(total_added_mtp_kv_layer_num) + layer_num = draft_config.get("num_hidden_layers", draft_config.get("n_layer")) + assert layer_num is not None, f"missing num_hidden_layers or n_layer in draft config: {draft_model_dir}" + return int(layer_num) @lru_cache(maxsize=None) diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 3d7274187b..930390171d 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -353,6 +353,27 @@ def get_config_dict(path): assert config_reads == ["/models/dspark"] +@pytest.mark.parametrize( + "mtp_mode, mtp_step, expected_layer_num", + [ + (None, 0, 0), + ("vanilla_no_att", 7, 0), + ("eagle_no_att", 7, 0), + ("vanilla_with_att", 7, 7), + ("eagle_with_att", 7, 1), + ], +) +def test_fixed_added_mtp_kv_layer_num_by_mode(monkeypatch, mtp_mode, mtp_step, expected_layer_num): + monkeypatch.setattr( + envs_utils, + "get_env_start_args", + lambda: SimpleNamespace(mtp_mode=mtp_mode, mtp_step=mtp_step), + ) + envs_utils.get_added_mtp_kv_layer_num.cache_clear() + + assert envs_utils.get_added_mtp_kv_layer_num() == expected_layer_num + + def test_dflash_added_kv_layers_come_from_draft_config(tmp_path): config_path = tmp_path / "config.json" config_path.write_text(json.dumps({"num_hidden_layers": 5})) From 637719e5df7413021815a48899a3fc4bdbe940d8 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 05:39:54 +0000 Subject: [PATCH 035/103] refactor: unify MTP model outputs --- lightllm/common/basemodel/basemodel.py | 39 ++++---- lightllm/common/basemodel/batch_objs.py | 77 +++++++++++++++- lightllm/common/basemodel/hidden_collector.py | 82 ++++++++++++----- lightllm/common/basemodel/mtp_manager.py | 8 +- .../common/basemodel/prefill_cuda_graph.py | 2 +- lightllm/models/qwen3_dspark/infer_struct.py | 8 -- .../layer_infer/post_layer_infer.py | 19 ++-- lightllm/models/qwen3_dspark/model.py | 33 ------- lightllm/models/qwen3_dspark/model_output.py | 20 ---- .../mode_backend/dp_backend/impl.py | 12 +-- .../mtp_speculative/proposers/dflash.py | 2 +- .../mtp_speculative/proposers/dspark.py | 8 +- .../mtp_speculative/proposers/eagle_mtp.py | 22 ++--- .../proposers/parallel_block.py | 2 +- .../mtp_speculative/proposers/vanilla_mtp.py | 15 +-- .../common/basemodel/test_hidden_collector.py | 39 ++++++-- .../common/basemodel/test_model_output.py | 12 +-- .../common/basemodel/test_mtp_manager.py | 9 +- .../models/test_qwen3_dspark_model_output.py | 92 ++++++++++++++----- .../mtp_speculative/test_eagle_overlap.py | 22 +++-- 20 files changed, 333 insertions(+), 190 deletions(-) delete mode 100644 lightllm/models/qwen3_dspark/infer_struct.py delete mode 100644 lightllm/models/qwen3_dspark/model_output.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index e62d2e4440..24188f1f7d 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -30,7 +30,6 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.hidden_collector import ( NoopHiddenCollector, - unpad_collected_hidden, ) from lightllm.common.basemodel.mtp_manager import MtpManager from lightllm.utils.custom_kernel_utis import pad2dim_tensor_to_new_batch @@ -547,7 +546,10 @@ def _create_unpad_decode_model_output(self, model_output: ModelOutput, origin_ba return model_output new_model_output = copy.copy(model_output) new_model_output.logits = new_model_output.logits[0:origin_batch_size] - new_model_output.spec_hidden = unpad_collected_hidden(new_model_output.spec_hidden, origin_batch_size) + new_model_output.collector = model_output.collector.unpad_decode( + padded_batch_size=padded_batch_size, + origin_batch_size=origin_batch_size, + ) return new_model_output def _create_unpad_prefill_model_output( @@ -556,7 +558,9 @@ def _create_unpad_prefill_model_output( new_model_output = copy.copy(padded_model_output) # logits 始终只对应每个请求最后一个位置,移除 padding 的 req 对应的行。 new_model_output.logits = new_model_output.logits[0:origin_batch_size] - new_model_output.spec_hidden = unpad_collected_hidden(new_model_output.spec_hidden, origin_handle_token_num) + new_model_output.collector = padded_model_output.collector.unpad_prefill( + origin_handle_token_num=origin_handle_token_num + ) # prompt_logics 保存整个 prefill 阶段所有 token 位置的 logits, # 按实际处理的 token 数量裁剪掉 padding 部分(仅 return_all_prompt_logics 模式下非空)。 if new_model_output.prompt_logics is not None: @@ -740,10 +744,9 @@ def prefill_func(input_tensors, _infer_state): predict_logits = self.post_infer.token_forward(last_input_embs, infer_state, self.pre_post_weight) hidden_collector = infer_state.hidden_collector hidden_collector.add_final_hidden(last_input_embs) - spec_hidden = hidden_collector.finish(infer_state=infer_state) model_output = ModelOutput( - logits=predict_logits, - spec_hidden=spec_hidden, + logits=predict_logits.contiguous(), + collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), prompt_logics=infer_state.prompt_logics, ) @@ -770,8 +773,10 @@ def _token_forward(self, infer_state: InferStateInfo): ) hidden_collector.add_final_hidden(last_input_embs) - spec_hidden = hidden_collector.finish(infer_state=infer_state) - model_output = ModelOutput(logits=predict_logits.contiguous(), spec_hidden=spec_hidden) + model_output = ModelOutput( + logits=predict_logits.contiguous(), + collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + ) # 在 cuda graph 模式下,输出需要转为 no ref tensor, 加强mem pool 的复用,降低显存的使用。 if infer_state.is_cuda_graph: @@ -1021,16 +1026,14 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state hidden_collector0.add_final_hidden(last_input_embs) hidden_collector1.add_final_hidden(last_input_embs1) - spec_hidden = hidden_collector0.finish(infer_state=infer_state) - spec_hidden1 = hidden_collector1.finish(infer_state=infer_state1) model_output = ModelOutput( logits=predict_logits.contiguous(), - spec_hidden=spec_hidden, + collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), prompt_logics=infer_state.prompt_logics, ) model_output1 = ModelOutput( logits=predict_logits1.contiguous(), - spec_hidden=spec_hidden1, + collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), prompt_logics=infer_state1.prompt_logics, ) @@ -1072,10 +1075,14 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: hidden_collector0.add_final_hidden(last_input_embs) hidden_collector1.add_final_hidden(last_input_embs1) - spec_hidden = hidden_collector0.finish(infer_state=infer_state) - spec_hidden1 = hidden_collector1.finish(infer_state=infer_state1) - model_output = ModelOutput(logits=predict_logits.contiguous(), spec_hidden=spec_hidden) - model_output1 = ModelOutput(logits=predict_logits1.contiguous(), spec_hidden=spec_hidden1) + model_output = ModelOutput( + logits=predict_logits.contiguous(), + collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + ) + model_output1 = ModelOutput( + logits=predict_logits1.contiguous(), + collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), + ) if infer_state.is_cuda_graph: model_output.to_no_ref_tensor() diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 0bfbc634b1..0be705a231 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -1,3 +1,4 @@ +import copy from dataclasses import dataclass from typing import List, Optional @@ -120,12 +121,77 @@ def _ensure_decode_group_metadata(self): self.b_mark_shared_group = torch.ones_like(self.b_req_idx, dtype=torch.int32) +@dataclass +class ModelMtpOutputCollector: + """保存一次模型 forward 为 MTP 推理产生的可选输出。""" + + # MTP drafter 使用的 hidden 特征。 + # - Vanilla MTP、EAGLE 主模型收集最终层 hidden;对应 draft 模型也会返回最终层 hidden, + # 供串行的下一层 MTP 模块或下一步自回归 draft 使用。 + # - EAGLE3、DFlash、DSpark 主模型收集 checkpoint 配置指定的若干 target layer hidden, + # 拼接后交给 draft 模型;EAGLE3 draft 模型还会返回最终层 hidden。 + # - DFlash、DSpark 的 block draft 模型不需要返回 hidden,因此该字段为 None。 + # - 未启用 MTP 时不收集投机特征,该字段同样为 None。 + spec_hidden: Optional[torch.Tensor] = None + + # DSpark block draft 模型直接生成的 token id,形状通常为 + # [request_count * block_size]。 + # - 仅 DSpark 启用 Markov head(markov_rank > 0)时由 head 直接生成并返回。 + # - DSpark 未启用 Markov head 时为 None,调用方从普通 logits 执行 argmax。 + # - Vanilla MTP、EAGLE、EAGLE3、DFlash 以及未启用 MTP 的模型均不使用该字段。 + draft_token_ids: Optional[torch.Tensor] = None + + # DSpark confidence head 输出的原始置信度 logits,形状通常为 + # [request_count, block_size],供动态 MTP verify 计算各 draft 位置的调度分数。 + # - 仅 DSpark checkpoint 启用 confidence head 时返回;动态 verify 模式要求该字段存在。 + # - 固定 verify 且未启用 confidence head 的 DSpark 模型可以返回 None。 + # - Vanilla MTP、EAGLE、EAGLE3、DFlash 以及未启用 MTP 的模型均不使用该字段。 + confidence_logits: Optional[torch.Tensor] = None + + def to_no_ref_tensor(self) -> None: + if self.spec_hidden is not None: + self.spec_hidden = tensor_to_no_ref_tensor(self.spec_hidden) + if self.draft_token_ids is not None: + self.draft_token_ids = tensor_to_no_ref_tensor(self.draft_token_ids) + if self.confidence_logits is not None: + self.confidence_logits = tensor_to_no_ref_tensor(self.confidence_logits) + + def unpad_decode( + self, padded_batch_size: int, origin_batch_size: int + ) -> "ModelMtpOutputCollector": + collector = copy.copy(self) + if collector.spec_hidden is not None: + collector.spec_hidden = collector.spec_hidden[:origin_batch_size] + if collector.draft_token_ids is not None: + collector.draft_token_ids = collector.draft_token_ids[:origin_batch_size] + if collector.confidence_logits is not None: + confidence_row_count = collector.confidence_logits.shape[0] + assert ( + confidence_row_count > 0 + and padded_batch_size % confidence_row_count == 0 + ) + rows_per_confidence = padded_batch_size // confidence_row_count + assert origin_batch_size % rows_per_confidence == 0 + collector.confidence_logits = collector.confidence_logits[ + : origin_batch_size // rows_per_confidence + ] + return collector + + def unpad_prefill(self, origin_handle_token_num: int) -> "ModelMtpOutputCollector": + collector = copy.copy(self) + if collector.spec_hidden is not None: + collector.spec_hidden = collector.spec_hidden[:origin_handle_token_num] + return collector + + @dataclass class ModelOutput: # 通用变量 logits: torch.Tensor - # Hidden states collected for the active speculative strategy. - spec_hidden: Optional[torch.Tensor] = None + # Collector is finalized by HiddenCollector.finish_output before being + # attached here. ModelOutput therefore owns a stable output view instead + # of the mutable collector used while the forward is still running. + collector: Optional[ModelMtpOutputCollector] = None # 用于判断 mem_indexes 是否成功写入 req manager 中的事件对象。 prefill_mem_indexes_ready_event: torch.Event = None @@ -135,7 +201,10 @@ class ModelOutput: # 需要返回 prompt logprobs 信息时才会非空。 prompt_logics: Optional[torch.Tensor] = None + def __post_init__(self) -> None: + if self.collector is None: + self.collector = ModelMtpOutputCollector() + def to_no_ref_tensor(self): self.logits = tensor_to_no_ref_tensor(self.logits) - if self.spec_hidden is not None: - self.spec_hidden = tensor_to_no_ref_tensor(self.spec_hidden) + self.collector.to_no_ref_tensor() diff --git a/lightllm/common/basemodel/hidden_collector.py b/lightllm/common/basemodel/hidden_collector.py index c8da1b0b08..3eb946fe82 100644 --- a/lightllm/common/basemodel/hidden_collector.py +++ b/lightllm/common/basemodel/hidden_collector.py @@ -7,21 +7,18 @@ import torch from transformers.configuration_utils import PretrainedConfig +from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor -def unpad_collected_hidden(hidden: Optional[torch.Tensor], token_count: int) -> Optional[torch.Tensor]: - return None if hidden is None else hidden[:token_count] - - class HiddenCollector(ABC): """Hidden state 收集器的抽象基类。 推理过程中,模型会在每一层计算完成后调用 :meth:`add`,并在一次 forward - 结束时调用 :meth:`finish` 生成供投机解码使用的 ``spec_hidden``。不同投机 - 解码模式可以通过子类决定不收集 hidden、只返回最终层 hidden,或者收集若干 - 中间层 hidden。 + 结束时调用 :meth:`finish_output` 生成统一的辅助输出。不同投机解码模式可以 + 通过子类决定不收集 hidden、只返回最终层 hidden、收集若干中间层 hidden, + 或同时收集 MTP head 生成的 token 与置信度。 BaseModel 持有一个不承载推理状态的 prototype,每个 InferStateInfo 通过 :meth:`new_instance` 获得独立收集器,使不同请求和 microbatch 的临时 hidden @@ -46,7 +43,7 @@ def restore_graph_state(self, graph_collector: "HiddenCollector") -> None: 基类默认不恢复任何数据。需要中间层 hidden 的子类应只复制保存 tensor 引用的容器,不复制 tensor 本身,使本次 infer state 可以读取 graph replay - 更新后的固定地址,同时在 :meth:`finish` 后独立清理自己的容器。 + 更新后的固定地址,同时在 :meth:`finish_output` 后独立清理自己的容器。 Args: graph_collector: capture 阶段保存在 Prefill CUDA Graph 中的只读 collector。 @@ -81,29 +78,42 @@ def add(self, layer_index: int, hidden: torch.Tensor) -> None: def add_final_hidden(self, final_hidden: torch.Tensor) -> None: """接收完成模型输出侧 gather 后的最终层 hidden tensor。 - BaseModel 在 logits 计算完成后、调用 :meth:`finish` 前调用该接口。基类 + BaseModel 在 logits 计算完成后、调用 :meth:`finish_output` 前调用该接口。基类 默认不保存 tensor;需要直接返回最终层 hidden 的子类应重写该方法并保存 - 引用,随后在 :meth:`finish` 中消费和清理。 + 引用,随后在 :meth:`finish_output` 中消费和清理。 Args: final_hidden: 已完成 TP/SP all-gather 和 DP unbalance 的最终层 hidden。 """ return + def add_mtp_outputs( + self, + draft_token_ids: Optional[torch.Tensor], + confidence_logits: Optional[torch.Tensor], + ) -> None: + """Collect optional token/confidence outputs produced by an MTP head. + + Only draft models with a specialized output head should call this + hook. Raising here makes an incorrect collector selection fail fast + instead of silently dropping model outputs. + """ + + raise RuntimeError(f"{self.__class__.__name__} does not collect MTP head outputs") + @abstractmethod - def finish(self, infer_state) -> Optional[torch.Tensor]: - """结束当前 microbatch 的收集并生成投机解码所需的 hidden tensor。 + def finish_output(self, infer_state) -> ModelMtpOutputCollector: + """结束当前 microbatch 的收集并生成统一的 MTP 输出。 子类应在此完成必要的拼接、TP/SP all-gather、DP unbalance 和 contiguous - 转换,并在返回前清理当前实例的临时状态。该接口在每次 forward 结束时 - 调用一次;不需要向 drafter 提供 hidden 的实现应返回 ``None``。 + 转换,将结果封装为 :class:`ModelMtpOutputCollector`,并在返回前清理当前 + 实例的临时状态。不需要提供额外 MTP 输出的实现应返回一个空 collector。 Args: infer_state: 当前 forward 的推理状态,包含通信拓扑、DP balance 等信息。 Returns: - 提供给投机解码 drafter 的连续 hidden tensor;当前模式不需要 hidden 时 - 返回 ``None``。 + 本次 forward 的 MTP 辅助输出;未启用 MTP 时返回内容为空的 collector。 """ raise NotImplementedError @@ -114,8 +124,36 @@ class NoopHiddenCollector(HiddenCollector): def new_instance(self) -> HiddenCollector: return NoopHiddenCollector() - def finish(self, infer_state) -> Optional[torch.Tensor]: - return None + def finish_output(self, infer_state) -> ModelMtpOutputCollector: + return ModelMtpOutputCollector() + + +class MtpHeadOutputCollector(NoopHiddenCollector): + """Collect outputs from an MTP draft head that does not expose hidden state.""" + + def __init__(self) -> None: + self.draft_token_ids: Optional[torch.Tensor] = None + self.confidence_logits: Optional[torch.Tensor] = None + + def new_instance(self) -> HiddenCollector: + return MtpHeadOutputCollector() + + def add_mtp_outputs( + self, + draft_token_ids: Optional[torch.Tensor], + confidence_logits: Optional[torch.Tensor], + ) -> None: + self.draft_token_ids = draft_token_ids + self.confidence_logits = confidence_logits + + def finish_output(self, infer_state) -> ModelMtpOutputCollector: + output = ModelMtpOutputCollector( + draft_token_ids=self.draft_token_ids, + confidence_logits=self.confidence_logits, + ) + self.draft_token_ids = None + self.confidence_logits = None + return output class FinalHiddenCollector(HiddenCollector): @@ -130,11 +168,11 @@ def new_instance(self) -> HiddenCollector: def add_final_hidden(self, final_hidden: torch.Tensor) -> None: self.final_hidden = final_hidden - def finish(self, infer_state) -> torch.Tensor: + def finish_output(self, infer_state) -> ModelMtpOutputCollector: assert self.final_hidden is not None final_hidden = self.final_hidden self.final_hidden = None - return final_hidden.contiguous() + return ModelMtpOutputCollector(spec_hidden=final_hidden.contiguous()) class LayerHiddenCollector(HiddenCollector): @@ -191,10 +229,10 @@ def _local_hidden(self) -> torch.Tensor: return self.layer_hiddens[0] return torch.cat(self.layer_hiddens, dim=-1) - def finish(self, infer_state) -> torch.Tensor: + def finish_output(self, infer_state) -> ModelMtpOutputCollector: local_hidden = self._local_hidden() self.layer_hiddens.clear() hidden = self.model.pre_infer._tpsp_allgather(input=local_hidden, infer_state=infer_state) if infer_state.need_dp_prefill_balance: hidden = infer_state._all_to_all_unbalance_get(data=hidden) - return hidden.contiguous() + return ModelMtpOutputCollector(spec_hidden=hidden.contiguous()) diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index 1d39f43e3f..be6c477b99 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -4,6 +4,7 @@ FinalHiddenCollector, HiddenCollector, LayerHiddenCollector, + MtpHeadOutputCollector, NoopHiddenCollector, ) from lightllm.utils.envs_utils import get_env_start_args @@ -82,7 +83,12 @@ def create_hidden_collector( if spec_mode is None: collector_type = NoopHiddenCollector elif model.is_mtp_draft_model: - collector_type = NoopHiddenCollector if spec_mode in self._BLOCK_DRAFT_MODES else FinalHiddenCollector + if spec_mode == "dspark": + collector_type = MtpHeadOutputCollector + elif spec_mode in self._BLOCK_DRAFT_MODES: + collector_type = NoopHiddenCollector + else: + collector_type = FinalHiddenCollector elif spec_mode in ("eagle3", *self._BLOCK_DRAFT_MODES): collector_type = LayerHiddenCollector collector_kwargs.update(model=model) diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index 6f8d91471d..5206ae5cdb 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -157,7 +157,7 @@ def _replay(self, input_tensors: List[torch.Tensor], infer_state: InferStateInfo graph_infer_state.copy_for_prefill_cuda_graph(new_infer_state=infer_state) # 首次 capture 后 replay 时,infer_state 与 graph_infer_state 是同一对象, - # 需要先创建运行时实例,避免 finish 清空 graph 中保存的 capture collector。 + # 需要先创建运行时实例,避免 finish_output 清空 graph 中保存的 capture collector。 if infer_state.hidden_collector is graph_hidden_collector: infer_state.hidden_collector = graph_hidden_collector.new_instance() infer_state.hidden_collector.restore_graph_state(graph_hidden_collector) diff --git a/lightllm/models/qwen3_dspark/infer_struct.py b/lightllm/models/qwen3_dspark/infer_struct.py deleted file mode 100644 index 0426bd87f7..0000000000 --- a/lightllm/models/qwen3_dspark/infer_struct.py +++ /dev/null @@ -1,8 +0,0 @@ -from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo - - -class Qwen3DSparkInferStateInfo(Qwen3DFlashInferStateInfo): - def __init__(self): - super().__init__() - self.confidence_logits = None - self.draft_token_ids = None diff --git a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py index 2cf407659c..5a74cd988e 100644 --- a/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py +++ b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py @@ -1,8 +1,8 @@ import torch from lightllm.distributed.communication_op import all_gather_into_tensor +from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo from lightllm.models.qwen3_dflash.layer_infer.post_layer_infer import Qwen3DFlashPostLayerInfer -from lightllm.models.qwen3_dspark.infer_struct import Qwen3DSparkInferStateInfo from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import ( Qwen3DSparkPreAndPostLayerWeight, ) @@ -84,7 +84,7 @@ def _sample_markov( self, local_logits: torch.Tensor, block_hidden: torch.Tensor, - infer_state: Qwen3DSparkInferStateInfo, + infer_state: Qwen3DFlashInferStateInfo, anchor_token_ids: torch.Tensor, layer_weight: Qwen3DSparkPreAndPostLayerWeight, ) -> torch.Tensor: @@ -140,7 +140,7 @@ def _sample_markov( def token_forward( self, input_embdings: torch.Tensor, - infer_state: Qwen3DSparkInferStateInfo, + infer_state: Qwen3DFlashInferStateInfo, layer_weight: Qwen3DSparkPreAndPostLayerWeight, ): if infer_state.is_prefill: @@ -166,23 +166,30 @@ def token_forward( anchor_token_ids=anchor_token_ids, layer_weight=layer_weight, ) - infer_state.draft_token_ids = sampled_tokens.reshape(-1) - infer_state.confidence_logits = self.predict_confidence_logits( + confidence_logits = self.predict_confidence_logits( block_hidden, anchor_token_ids=anchor_token_ids, sampled_tokens=sampled_tokens, layer_weight=layer_weight, ) + infer_state.hidden_collector.add_mtp_outputs( + draft_token_ids=sampled_tokens.reshape(-1), + confidence_logits=confidence_logits, + ) # Graph unpadding still uses the leading logits dimension when token ids are returned directly. return local_logits.new_empty((token_num, 1)) logits = self._lm_head_and_gather(last_input, token_num, layer_weight, infer_state) block_logits = logits.reshape(num_reqs, self.block_size_, -1) sampled_tokens = torch.argmax(block_logits, dim=-1) - infer_state.confidence_logits = self.predict_confidence_logits( + confidence_logits = self.predict_confidence_logits( block_hidden, anchor_token_ids=anchor_token_ids, sampled_tokens=sampled_tokens, layer_weight=layer_weight, ) + infer_state.hidden_collector.add_mtp_outputs( + draft_token_ids=None, + confidence_logits=confidence_logits, + ) return logits diff --git a/lightllm/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py index 9e42d02d41..c857248c98 100644 --- a/lightllm/models/qwen3_dspark/model.py +++ b/lightllm/models/qwen3_dspark/model.py @@ -1,10 +1,7 @@ from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel from lightllm.models.draft_registry import DraftModelRegistry -from lightllm.models.qwen3_dspark.infer_struct import Qwen3DSparkInferStateInfo from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer -from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import Qwen3DSparkPreAndPostLayerWeight -from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor @DraftModelRegistry(model_type="qwen3", spec_modes="dspark") @@ -18,33 +15,3 @@ class Qwen3DSparkModel(Qwen3DFlashModel): pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight post_layer_infer_class = Qwen3DSparkPostLayerInfer - infer_state_class = Qwen3DSparkInferStateInfo - - def _token_forward(self, infer_state: Qwen3DSparkInferStateInfo): - model_output = super()._token_forward(infer_state) - if infer_state.is_cuda_graph: - if infer_state.confidence_logits is not None: - infer_state.confidence_logits = tensor_to_no_ref_tensor(infer_state.confidence_logits) - if infer_state.draft_token_ids is not None: - infer_state.draft_token_ids = tensor_to_no_ref_tensor(infer_state.draft_token_ids) - - return DSparkModelOutput( - logits=model_output.logits, - spec_hidden=model_output.spec_hidden, - confidence_logits=infer_state.confidence_logits, - draft_token_ids=infer_state.draft_token_ids, - ) - - def _create_unpad_decode_model_output(self, model_output: DSparkModelOutput, origin_batch_size: int): - padded_batch_size = model_output.logits.shape[0] - model_output = super()._create_unpad_decode_model_output(model_output, origin_batch_size) - if padded_batch_size == origin_batch_size: - return model_output - - if model_output.draft_token_ids is not None: - model_output.draft_token_ids = model_output.draft_token_ids[:origin_batch_size] - if model_output.confidence_logits is not None: - confidence_rows = model_output.confidence_logits.shape[0] - rows_per_confidence = padded_batch_size // confidence_rows - model_output.confidence_logits = model_output.confidence_logits[: origin_batch_size // rows_per_confidence] - return model_output diff --git a/lightllm/models/qwen3_dspark/model_output.py b/lightllm/models/qwen3_dspark/model_output.py deleted file mode 100644 index bff766ec21..0000000000 --- a/lightllm/models/qwen3_dspark/model_output.py +++ /dev/null @@ -1,20 +0,0 @@ -from dataclasses import dataclass -from typing import Optional - -import torch - -from lightllm.common.basemodel.batch_objs import ModelOutput -from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor - - -@dataclass -class DSparkModelOutput(ModelOutput): - confidence_logits: Optional[torch.Tensor] = None - draft_token_ids: Optional[torch.Tensor] = None - - def to_no_ref_tensor(self): - super().to_no_ref_tensor() - if self.confidence_logits is not None: - self.confidence_logits = tensor_to_no_ref_tensor(self.confidence_logits) - if self.draft_token_ids is not None: - self.draft_token_ids = tensor_to_no_ref_tensor(self.draft_token_ids) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 054e84a6c2..fb0c078d02 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -614,7 +614,7 @@ def _draft_decode_vanilla( all_next_token_ids = [] # share some inference info with the main model draft_model_input = model_input - draft_hidden = model_output.spec_hidden + draft_hidden = model_output.collector.spec_hidden draft_next_token_ids_gpu = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input.batch_size, @@ -630,7 +630,7 @@ def _draft_decode_vanilla( draft_model_input.mtp_draft_input_hiddens = draft_hidden # spec decode: MTP draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) - draft_hidden = draft_model_output.spec_hidden + draft_hidden = draft_model_output.collector.spec_hidden draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) all_next_token_ids.append(draft_next_token_ids_gpu) @@ -946,8 +946,8 @@ def _draft_decode_vanilla_overlap( all_next_token_ids.append(next_token_ids) # share some inference info with the main model draft_model_input0, draft_model_input1 = model_input0, model_input1 - draft_hidden0 = model_output0.spec_hidden - draft_hidden1 = model_output1.spec_hidden + draft_hidden0 = model_output0.collector.spec_hidden + draft_hidden1 = model_output1.collector.spec_hidden draft_next_token_ids_gpu0 = self._build_padded_next_token_ids( token_ids=next_token_ids, @@ -974,8 +974,8 @@ def _draft_decode_vanilla_overlap( draft_model_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( draft_model_input0, draft_model_input1 ) - draft_hidden0 = draft_model_output0.spec_hidden - draft_hidden1 = draft_model_output1.spec_hidden + draft_hidden0 = draft_model_output0.collector.spec_hidden + draft_hidden1 = draft_model_output1.collector.spec_hidden draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index 07dbcfbfe9..522aec6c17 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -44,7 +44,7 @@ def propose_next( ) self.extend_draft_kv_cache( main_model_input=main_model_input, - target_hidden=main_model_output.spec_hidden, + target_hidden=main_model_output.collector.spec_hidden, ) draft_output = draft_model.forward(draft_input) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 2e981fc7a3..bbe2da75e6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -54,19 +54,19 @@ def propose_next( ) self.extend_draft_kv_cache( main_model_input=main_model_input, - target_hidden=main_model_output.spec_hidden, + target_hidden=main_model_output.collector.spec_hidden, ) draft_output = draft_model.forward(draft_input) - if draft_output.draft_token_ids is None: + if draft_output.collector.draft_token_ids is None: flat_draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) else: - flat_draft_token_ids = draft_output.draft_token_ids + flat_draft_token_ids = draft_output.collector.draft_token_ids block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) proposal_token_ids[accepted_tail_rows, 1:] = block_draft_token_ids[:, :draft_step] if self.enable_dynmaic_mtp: - confidence_logits = draft_output.confidence_logits + confidence_logits = draft_output.collector.confidence_logits if confidence_logits is None: raise RuntimeError("DSpark dynamic verify requires confidence head logits") # Match the clamp used by the GPU dynamic row selector before it diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py index 5245ec2cd1..82a3c5989a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py @@ -37,7 +37,7 @@ def build_draft_state_from_prefill( prepare_mtp_prefill_inputs( model_input=target_model_input, b_next_token_ids=next_token_ids, - mtp_draft_input_hiddens=target_model_output.spec_hidden, + mtp_draft_input_hiddens=target_model_output.collector.spec_hidden, ) self.backend.draft_models[0].forward(target_model_input) @@ -55,12 +55,12 @@ def build_draft_state_from_prefill_overlap( prepare_mtp_prefill_inputs( model_input=target_model_input0, b_next_token_ids=next_token_ids0, - mtp_draft_input_hiddens=target_model_output0.spec_hidden, + mtp_draft_input_hiddens=target_model_output0.collector.spec_hidden, ) prepare_mtp_prefill_inputs( model_input=target_model_input1, b_next_token_ids=next_token_ids1, - mtp_draft_input_hiddens=target_model_output1.spec_hidden, + mtp_draft_input_hiddens=target_model_output1.collector.spec_hidden, ) self.backend.draft_models[0].microbatch_overlap_prefill(target_model_input0, target_model_input1) @@ -143,7 +143,7 @@ def propose_next( self.prepare_verify_extend_input( model_input=main_model_input, input_ids=next_token_ids, - target_hidden=main_model_output.spec_hidden, + target_hidden=main_model_output.collector.spec_hidden, ) extend_output = draft_model.forward(main_model_input) @@ -154,7 +154,7 @@ def propose_next( else: draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) proposal_token_ids[accepted_tail_rows, 1] = draft_token_ids - draft_hidden = extend_output.spec_hidden.index_select(0, accepted_tail_rows) + draft_hidden = extend_output.collector.spec_hidden.index_select(0, accepted_tail_rows) if draft_step == 1: return SpecProposal( @@ -196,7 +196,7 @@ def propose_next( else: draft_token_ids = self._gen_argmax_token_ids(draft_output) proposal_token_ids[accepted_tail_rows, step + 1] = draft_token_ids - draft_hidden = draft_output.spec_hidden + draft_hidden = draft_output.collector.spec_hidden draft_seq_lens.add_(1) return SpecProposal( @@ -262,7 +262,7 @@ def propose_next_overlap( self.prepare_verify_extend_input( model_input=model_input, input_ids=token_ids, - target_hidden=model_output.spec_hidden, + target_hidden=model_output.collector.spec_hidden, ) verify_row_count = real_verify_rows0 + real_verify_rows1 @@ -288,7 +288,7 @@ def propose_next_overlap( accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) draft_token_ids_by_batch.append(draft_token_ids) - draft_hiddens_by_batch.append(extend_output.spec_hidden.index_select(0, accepted_tail_rows)) + draft_hiddens_by_batch.append(extend_output.collector.spec_hidden.index_select(0, accepted_tail_rows)) draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) proposal_rows = accepted_tail_rows[:real_request_count] + proposal_row_offsets[batch_index] @@ -346,7 +346,7 @@ def propose_next_overlap( for batch_index, draft_output in enumerate(draft_outputs): draft_token_ids = self._gen_argmax_token_ids(draft_output) draft_token_ids_by_batch[batch_index] = draft_token_ids - draft_hiddens_by_batch[batch_index] = draft_output.spec_hidden + draft_hiddens_by_batch[batch_index] = draft_output.collector.spec_hidden draft_seq_lens_by_batch[batch_index].add_(1) real_request_count = real_request_counts[batch_index] proposal_token_ids[proposal_rows_by_batch[batch_index], step + 1] = draft_token_ids[:real_request_count] @@ -406,7 +406,7 @@ def propose_next_overlap( ) draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] - draft_hiddens_by_batch = [main_model_output0.spec_hidden, main_model_output1.spec_hidden] + draft_hiddens_by_batch = [main_model_output0.collector.spec_hidden, main_model_output1.collector.spec_hidden] draft_model = self.backend.draft_models[0] hold_mem_index = self.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX @@ -438,7 +438,7 @@ def propose_next_overlap( ).view(-1) draft_token_ids_by_batch[batch_index] = self._gen_argmax_token_ids(draft_output) - draft_hiddens_by_batch[batch_index] = draft_output.spec_hidden + draft_hiddens_by_batch[batch_index] = draft_output.collector.spec_hidden proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids_by_batch[0][:real_verify_rows0] proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids_by_batch[1][:real_verify_rows1] diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py index 1b345f01cc..b60ab20ea7 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py @@ -40,7 +40,7 @@ def build_draft_state_from_prefill( target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - target_hidden = target_model_output.spec_hidden + target_hidden = target_model_output.collector.spec_hidden if target_hidden.numel() == 0: return diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py index a8176392c8..aa0b94f49d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py @@ -34,7 +34,7 @@ def build_draft_state_from_prefill( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - draft_hidden = target_model_output.spec_hidden + draft_hidden = target_model_output.collector.spec_hidden draft_token_ids = next_token_ids for draft_model in self.backend.draft_models: prepare_mtp_prefill_inputs( @@ -43,7 +43,7 @@ def build_draft_state_from_prefill( mtp_draft_input_hiddens=draft_hidden, ) draft_output = draft_model.forward(target_model_input) - draft_hidden = draft_output.spec_hidden + draft_hidden = draft_output.collector.spec_hidden draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) def build_draft_state_from_prefill_overlap( @@ -57,7 +57,10 @@ def build_draft_state_from_prefill_overlap( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - draft_hiddens_by_batch = [target_model_output0.spec_hidden, target_model_output1.spec_hidden] + draft_hiddens_by_batch = [ + target_model_output0.collector.spec_hidden, + target_model_output1.collector.spec_hidden, + ] draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] for draft_model in self.backend.draft_models: @@ -76,7 +79,7 @@ def build_draft_state_from_prefill_overlap( target_model_input1, ) for batch_index, draft_output in enumerate(draft_outputs): - draft_hiddens_by_batch[batch_index] = draft_output.spec_hidden + draft_hiddens_by_batch[batch_index] = draft_output.collector.spec_hidden draft_token_ids_by_batch[batch_index] = self.backend._gen_argmax_token_ids(draft_output) def propose_next( @@ -90,7 +93,7 @@ def propose_next( ) -> SpecProposal: verify_row_count = int(next_token_ids.shape[0]) draft_token_ids = next_token_ids - draft_hidden = main_model_output.spec_hidden + draft_hidden = main_model_output.collector.spec_hidden proposal_token_ids = next_token_ids.new_empty((verify_row_count, draft_step + 1)) proposal_token_ids[:, 0] = next_token_ids schedule_scores = ( @@ -108,7 +111,7 @@ def propose_next( main_model_input.input_ids = draft_token_ids main_model_input.mtp_draft_input_hiddens = draft_hidden draft_output = draft_model.forward(main_model_input) - draft_hidden = draft_output.spec_hidden + draft_hidden = draft_output.collector.spec_hidden if self.enable_dynmaic_mtp: draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) schedule_scores[:, step] = draft_token_probs diff --git a/unit_tests/common/basemodel/test_hidden_collector.py b/unit_tests/common/basemodel/test_hidden_collector.py index 3ce9d881ca..22748b7294 100644 --- a/unit_tests/common/basemodel/test_hidden_collector.py +++ b/unit_tests/common/basemodel/test_hidden_collector.py @@ -9,6 +9,7 @@ FinalHiddenCollector, HiddenCollector, LayerHiddenCollector, + MtpHeadOutputCollector, NoopHiddenCollector, ) @@ -49,14 +50,14 @@ def test_final_hidden_collectors_are_independent_instances(): infer_state = SimpleNamespace(need_dp_prefill_balance=False) collector0.add_final_hidden(hidden0) - collected = collector0.finish(infer_state=infer_state) + collected = collector0.finish_output(infer_state=infer_state).spec_hidden assert collected.data_ptr() == hidden0.data_ptr() assert collector0.final_hidden is None collector0.add_final_hidden(hidden0) collector1.add_final_hidden(hidden1) - collected0 = collector0.finish(infer_state=infer_state) - collected1 = collector1.finish(infer_state=infer_state) + collected0 = collector0.finish_output(infer_state=infer_state).spec_hidden + collected1 = collector1.finish_output(infer_state=infer_state).spec_hidden assert collected0.data_ptr() == hidden0.data_ptr() assert collected1.data_ptr() == hidden1.data_ptr() @@ -74,8 +75,8 @@ def test_layer_hidden_collector_keeps_microbatch_state_separate(monkeypatch): collector0.add(layer_index=0, hidden=hidden0) collector1.add(layer_index=0, hidden=hidden1) - collected0 = collector0.finish(infer_state=infer_state) - collected1 = collector1.finish(infer_state=infer_state) + collected0 = collector0.finish_output(infer_state=infer_state).spec_hidden + collected1 = collector1.finish_output(infer_state=infer_state).spec_hidden assert torch.equal(collected0, hidden0) assert torch.equal(collected1, hidden1) @@ -96,14 +97,32 @@ def test_noop_collector_keeps_normal_forward_output_minimal(): collector.add(layer_index=0, hidden=final_hidden) collector.add_final_hidden(final_hidden) - assert collector.finish(infer_state=None) is None + assert collector.finish_output(infer_state=None).spec_hidden is None + + +def test_mtp_head_output_collector_returns_and_clears_outputs(): + collector = MtpHeadOutputCollector() + draft_token_ids = torch.arange(6) + confidence_logits = torch.arange(6).view(2, 3) + + collector.add_mtp_outputs( + draft_token_ids=draft_token_ids, + confidence_logits=confidence_logits, + ) + output = collector.finish_output(infer_state=None) + + assert output.spec_hidden is None + assert output.draft_token_ids is draft_token_ids + assert output.confidence_logits is confidence_logits + assert collector.draft_token_ids is None + assert collector.confidence_logits is None def test_final_collector_returns_final_hidden_without_layer_bookkeeping(): final_hidden = torch.randn(2, 3) collector = FinalHiddenCollector() collector.add_final_hidden(final_hidden) - collected = collector.finish(infer_state=None) + collected = collector.finish_output(infer_state=None).spec_hidden assert collected.data_ptr() == final_hidden.data_ptr() @@ -121,14 +140,14 @@ def test_layer_collector_preserves_selected_layers_in_model_order(monkeypatch): collector.add(layer_index=2, hidden=layer2) layer0.fill_(9.0) - collected = collector.finish(infer_state=SimpleNamespace(need_dp_prefill_balance=False)) + collected = collector.finish_output(infer_state=SimpleNamespace(need_dp_prefill_balance=False)).spec_hidden assert torch.equal(collected, torch.cat([torch.full((2, 2), 1.0), layer2], dim=-1)) assert not collector.layer_hiddens collector.add(layer_index=0, hidden=layer0) collector.add(layer_index=2, hidden=layer2) - collected = collector.finish(infer_state=SimpleNamespace(need_dp_prefill_balance=False)) + collected = collector.finish_output(infer_state=SimpleNamespace(need_dp_prefill_balance=False)).spec_hidden assert torch.equal(collected, torch.cat([layer0, layer2], dim=-1)) assert not collector.layer_hiddens @@ -143,7 +162,7 @@ def test_layer_collector_restores_graph_state_without_sharing_runtime_container( graph_collector.add(layer_index=0, hidden=torch.full((2, 3), 1.0)) collector.restore_graph_state(graph_collector) - collected = collector.finish(infer_state=infer_state) + collected = collector.finish_output(infer_state=infer_state).spec_hidden assert torch.equal(collected, torch.full((2, 3), 1.0)) assert not collector.layer_hiddens diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index e66c120123..4f09cf7099 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -1,31 +1,31 @@ import torch from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput def test_decode_unpad_slices_spec_output_with_logits(): model = TpPartBaseModel.__new__(TpPartBaseModel) output = ModelOutput( logits=torch.arange(24).view(6, 4), - spec_hidden=torch.arange(18).view(6, 3), + collector=ModelMtpOutputCollector(spec_hidden=torch.arange(18).view(6, 3)), ) unpadded = model._create_unpad_decode_model_output(output, origin_batch_size=4) assert unpadded.logits.shape == (4, 4) - assert unpadded.spec_hidden.shape == (4, 3) + assert unpadded.collector.spec_hidden.shape == (4, 3) # Unpadding returns a shallow output copy and leaves the graph-owned # tensors on the original ModelOutput intact. assert output.logits.shape == (6, 4) - assert output.spec_hidden.shape == (6, 3) + assert output.collector.spec_hidden.shape == (6, 3) def test_prefill_unpad_uses_token_rows_for_spec_hidden(): model = TpPartBaseModel.__new__(TpPartBaseModel) output = ModelOutput( logits=torch.arange(20).view(5, 4), - spec_hidden=torch.arange(24).view(8, 3), + collector=ModelMtpOutputCollector(spec_hidden=torch.arange(24).view(8, 3)), prompt_logics=torch.arange(32).view(8, 4), ) @@ -36,5 +36,5 @@ def test_prefill_unpad_uses_token_rows_for_spec_hidden(): ) assert unpadded.logits.shape == (3, 4) - assert unpadded.spec_hidden.shape == (6, 3) + assert unpadded.collector.spec_hidden.shape == (6, 3) assert unpadded.prompt_logics.shape == (6, 4) diff --git a/unit_tests/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py index ce622000cd..c7ad04f56b 100644 --- a/unit_tests/common/basemodel/test_mtp_manager.py +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -4,7 +4,12 @@ import lightllm.common.basemodel.hidden_collector as hidden_collector_module import lightllm.common.basemodel.mtp_manager as mtp_manager_module -from lightllm.common.basemodel.hidden_collector import FinalHiddenCollector, LayerHiddenCollector, NoopHiddenCollector +from lightllm.common.basemodel.hidden_collector import ( + FinalHiddenCollector, + LayerHiddenCollector, + MtpHeadOutputCollector, + NoopHiddenCollector, +) from lightllm.common.basemodel.mtp_manager import MtpManager @@ -96,7 +101,7 @@ def test_get_instance_returns_singleton(monkeypatch): ("eagle3", False, LayerHiddenCollector), ("dspark", False, LayerHiddenCollector), ("eagle3", True, FinalHiddenCollector), - ("dspark", True, NoopHiddenCollector), + ("dspark", True, MtpHeadOutputCollector), ], ) def test_create_hidden_collector_selects_implementation(monkeypatch, spec_mode, is_draft_model, expected_type): diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index f621e6f7bf..f967e86ffb 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -4,64 +4,106 @@ import torch from lightllm.common.basemodel import batch_objs +from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel -from lightllm.models.qwen3_dspark import model_output as dspark_model_output from lightllm.models.qwen3_dspark.layer_weights import pre_and_post_layer_weight as dspark_pre_post_weight from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel -from lightllm.models.qwen3_dspark.model_output import DSparkModelOutput @pytest.mark.parametrize("model_class", [Qwen3DSparkModel, Qwen3_5DSparkModel]) -def test_dspark_decode_unpad_preserves_output_type_and_slices_dspark_fields(model_class): +def test_dspark_decode_unpad_uses_common_output_and_slices_mtp_fields(model_class): model = model_class.__new__(model_class) - output = DSparkModelOutput( + output = ModelOutput( logits=torch.arange(48).view(12, 4), - spec_hidden=torch.arange(36).view(12, 3), - confidence_logits=torch.arange(12).view(3, 4), - draft_token_ids=torch.arange(12), + collector=ModelMtpOutputCollector( + spec_hidden=torch.arange(36).view(12, 3), + confidence_logits=torch.arange(12).view(3, 4), + draft_token_ids=torch.arange(12), + ), ) unpadded = model._create_unpad_decode_model_output(output, origin_batch_size=8) - assert isinstance(unpadded, DSparkModelOutput) + assert isinstance(unpadded, ModelOutput) assert unpadded.logits.shape == (8, 4) - assert unpadded.spec_hidden.shape == (8, 3) - assert unpadded.confidence_logits.shape == (2, 4) - assert unpadded.draft_token_ids.shape == (8,) + assert unpadded.collector.spec_hidden.shape == (8, 3) + assert unpadded.collector.confidence_logits.shape == (2, 4) + assert unpadded.collector.draft_token_ids.shape == (8,) assert output.logits.shape == (12, 4) - assert output.confidence_logits.shape == (3, 4) - assert output.draft_token_ids.shape == (12,) + assert output.collector.confidence_logits.shape == (3, 4) + assert output.collector.draft_token_ids.shape == (12,) -def test_dspark_no_ref_conversion_dispatches_to_dspark_fields(monkeypatch): +def test_common_output_no_ref_conversion_includes_dspark_fields(monkeypatch): monkeypatch.setattr(batch_objs, "tensor_to_no_ref_tensor", torch.clone) - monkeypatch.setattr(dspark_model_output, "tensor_to_no_ref_tensor", torch.clone) - output = DSparkModelOutput( + output = ModelOutput( logits=torch.ones((2, 4)), - spec_hidden=torch.ones((2, 3)), - confidence_logits=torch.ones((1, 2)), - draft_token_ids=torch.ones((2,), dtype=torch.int64), + collector=ModelMtpOutputCollector( + spec_hidden=torch.ones((2, 3)), + confidence_logits=torch.ones((1, 2)), + draft_token_ids=torch.ones((2,), dtype=torch.int64), + ), ) original_ptrs = ( output.logits.data_ptr(), - output.spec_hidden.data_ptr(), - output.confidence_logits.data_ptr(), - output.draft_token_ids.data_ptr(), + output.collector.spec_hidden.data_ptr(), + output.collector.confidence_logits.data_ptr(), + output.collector.draft_token_ids.data_ptr(), ) output.to_no_ref_tensor() converted_ptrs = ( output.logits.data_ptr(), - output.spec_hidden.data_ptr(), - output.confidence_logits.data_ptr(), - output.draft_token_ids.data_ptr(), + output.collector.spec_hidden.data_ptr(), + output.collector.confidence_logits.data_ptr(), + output.collector.draft_token_ids.data_ptr(), ) assert all(converted != original for converted, original in zip(converted_ptrs, original_ptrs)) +def test_dspark_post_layer_publishes_head_results_through_collector(): + post_infer = Qwen3DSparkPostLayerInfer.__new__(Qwen3DSparkPostLayerInfer) + post_infer.block_size_ = 2 + post_infer.markov_rank_ = 2 + post_infer._slice_get_last_input = lambda input_embeddings, infer_state: ( + input_embeddings, + 4, + ) + post_infer._norm = lambda hidden, infer_state, layer_weight: hidden + post_infer._sample_markov = lambda *args, **kwargs: torch.tensor([[1, 2], [3, 4]]) + confidence_logits = torch.arange(4, dtype=torch.float32).view(2, 2) + post_infer.predict_confidence_logits = lambda *args, **kwargs: confidence_logits + + class RecordingCollector: + def add_mtp_outputs(self, **kwargs): + self.outputs = kwargs + + class LMHead: + def __call__(self, input, alloc_func): + return torch.empty((8, input.shape[1])) + + collector = RecordingCollector() + infer_state = SimpleNamespace( + is_prefill=False, + input_ids=torch.tensor([10, 0, 20, 0]), + hidden_collector=collector, + ) + layer_weight = SimpleNamespace(lm_head_weight_=LMHead()) + + logits = post_infer.token_forward( + input_embdings=torch.randn(4, 3), + infer_state=infer_state, + layer_weight=layer_weight, + ) + + assert logits.shape == (4, 1) + assert torch.equal(collector.outputs["draft_token_ids"], torch.tensor([1, 2, 3, 4])) + assert collector.outputs["confidence_logits"] is confidence_logits + + @pytest.mark.parametrize( "head_config", [ diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 9ad007fb4b..bb7846ab45 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -2,7 +2,7 @@ import torch -from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import ( AutoregressiveEagleProposer, @@ -23,7 +23,7 @@ def microbatch_overlap_prefill(self, input0, input1): return tuple( ModelOutput( logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), - spec_hidden=torch.ones((model_input.batch_size, 2)), + collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), ) for model_input in (input0, input1) ) @@ -34,7 +34,7 @@ def microbatch_overlap_decode(self, input0, input1): return tuple( ModelOutput( logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), - spec_hidden=torch.ones((model_input.batch_size, 2)), + collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), ) for model_input in (input0, input1) ) @@ -78,12 +78,16 @@ def test_overlap_eagle_keeps_fixed_verify_layout(): proposal = proposer.propose_next_overlap( main_model_input0=model_input0, - main_model_output0=ModelOutput(logits=torch.empty((6, 1)), spec_hidden=torch.ones((6, 2))), + main_model_output0=ModelOutput( + logits=torch.empty((6, 1)), collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + ), next_token_ids0=torch.arange(6, dtype=torch.int64), real_verify_rows0=3, accept_len0=torch.tensor([2, 1], dtype=torch.int32), main_model_input1=model_input1, - main_model_output1=ModelOutput(logits=torch.empty((6, 1)), spec_hidden=torch.ones((6, 2))), + main_model_output1=ModelOutput( + logits=torch.empty((6, 1)), collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + ), next_token_ids1=torch.arange(10, 16, dtype=torch.int64), real_verify_rows1=6, accept_len1=torch.tensor([1, 3], dtype=torch.int32), @@ -121,12 +125,16 @@ def test_autoregressive_eagle_reuses_overlap_inputs(): proposal = proposer.propose_next_overlap( main_model_input0=model_input0, - main_model_output0=ModelOutput(logits=torch.empty((6, 1)), spec_hidden=torch.ones((6, 2))), + main_model_output0=ModelOutput( + logits=torch.empty((6, 1)), collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + ), next_token_ids0=torch.arange(6, dtype=torch.int64), real_verify_rows0=3, accept_len0=torch.tensor([2, 1], dtype=torch.int32), main_model_input1=model_input1, - main_model_output1=ModelOutput(logits=torch.empty((6, 1)), spec_hidden=torch.ones((6, 2))), + main_model_output1=ModelOutput( + logits=torch.empty((6, 1)), collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + ), next_token_ids1=torch.arange(10, 16, dtype=torch.int64), real_verify_rows1=6, accept_len1=torch.tensor([1, 3], dtype=torch.int32), From 76faef97e9bf76c1809f1734f2272fb2ac8cbc9e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 05:55:22 +0000 Subject: [PATCH 036/103] refactor: clarify MTP output collector naming --- lightllm/common/basemodel/basemodel.py | 16 ++++++------ lightllm/common/basemodel/batch_objs.py | 10 +++---- .../mode_backend/dp_backend/impl.py | 12 ++++----- .../mtp_speculative/proposers/dflash.py | 2 +- .../mtp_speculative/proposers/dspark.py | 8 +++--- .../mtp_speculative/proposers/eagle_mtp.py | 25 ++++++++++-------- .../proposers/parallel_block.py | 2 +- .../mtp_speculative/proposers/vanilla_mtp.py | 14 +++++----- .../common/basemodel/test_model_output.py | 10 +++---- .../models/test_qwen3_dspark_model_output.py | 26 +++++++++---------- .../mtp_speculative/test_eagle_overlap.py | 12 ++++----- 11 files changed, 70 insertions(+), 67 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 24188f1f7d..bce508adb8 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -546,7 +546,7 @@ def _create_unpad_decode_model_output(self, model_output: ModelOutput, origin_ba return model_output new_model_output = copy.copy(model_output) new_model_output.logits = new_model_output.logits[0:origin_batch_size] - new_model_output.collector = model_output.collector.unpad_decode( + new_model_output.mtp_collector = model_output.mtp_collector.unpad_decode( padded_batch_size=padded_batch_size, origin_batch_size=origin_batch_size, ) @@ -558,7 +558,7 @@ def _create_unpad_prefill_model_output( new_model_output = copy.copy(padded_model_output) # logits 始终只对应每个请求最后一个位置,移除 padding 的 req 对应的行。 new_model_output.logits = new_model_output.logits[0:origin_batch_size] - new_model_output.collector = padded_model_output.collector.unpad_prefill( + new_model_output.mtp_collector = padded_model_output.mtp_collector.unpad_prefill( origin_handle_token_num=origin_handle_token_num ) # prompt_logics 保存整个 prefill 阶段所有 token 位置的 logits, @@ -746,7 +746,7 @@ def prefill_func(input_tensors, _infer_state): hidden_collector.add_final_hidden(last_input_embs) model_output = ModelOutput( logits=predict_logits.contiguous(), - collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), prompt_logics=infer_state.prompt_logics, ) @@ -775,7 +775,7 @@ def _token_forward(self, infer_state: InferStateInfo): hidden_collector.add_final_hidden(last_input_embs) model_output = ModelOutput( logits=predict_logits.contiguous(), - collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), ) # 在 cuda graph 模式下,输出需要转为 no ref tensor, 加强mem pool 的复用,降低显存的使用。 @@ -1028,12 +1028,12 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state hidden_collector1.add_final_hidden(last_input_embs1) model_output = ModelOutput( logits=predict_logits.contiguous(), - collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), prompt_logics=infer_state.prompt_logics, ) model_output1 = ModelOutput( logits=predict_logits1.contiguous(), - collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), + mtp_collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), prompt_logics=infer_state1.prompt_logics, ) @@ -1077,11 +1077,11 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: hidden_collector1.add_final_hidden(last_input_embs1) model_output = ModelOutput( logits=predict_logits.contiguous(), - collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), ) model_output1 = ModelOutput( logits=predict_logits1.contiguous(), - collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), + mtp_collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), ) if infer_state.is_cuda_graph: diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 0be705a231..169edd6661 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -188,10 +188,10 @@ def unpad_prefill(self, origin_handle_token_num: int) -> "ModelMtpOutputCollecto class ModelOutput: # 通用变量 logits: torch.Tensor - # Collector is finalized by HiddenCollector.finish_output before being + # MTP collector is finalized by HiddenCollector.finish_output before being # attached here. ModelOutput therefore owns a stable output view instead # of the mutable collector used while the forward is still running. - collector: Optional[ModelMtpOutputCollector] = None + mtp_collector: Optional[ModelMtpOutputCollector] = None # 用于判断 mem_indexes 是否成功写入 req manager 中的事件对象。 prefill_mem_indexes_ready_event: torch.Event = None @@ -202,9 +202,9 @@ class ModelOutput: prompt_logics: Optional[torch.Tensor] = None def __post_init__(self) -> None: - if self.collector is None: - self.collector = ModelMtpOutputCollector() + if self.mtp_collector is None: + self.mtp_collector = ModelMtpOutputCollector() def to_no_ref_tensor(self): self.logits = tensor_to_no_ref_tensor(self.logits) - self.collector.to_no_ref_tensor() + self.mtp_collector.to_no_ref_tensor() diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index fb0c078d02..fed67950e0 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -614,7 +614,7 @@ def _draft_decode_vanilla( all_next_token_ids = [] # share some inference info with the main model draft_model_input = model_input - draft_hidden = model_output.collector.spec_hidden + draft_hidden = model_output.mtp_collector.spec_hidden draft_next_token_ids_gpu = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input.batch_size, @@ -630,7 +630,7 @@ def _draft_decode_vanilla( draft_model_input.mtp_draft_input_hiddens = draft_hidden # spec decode: MTP draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) - draft_hidden = draft_model_output.collector.spec_hidden + draft_hidden = draft_model_output.mtp_collector.spec_hidden draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) all_next_token_ids.append(draft_next_token_ids_gpu) @@ -946,8 +946,8 @@ def _draft_decode_vanilla_overlap( all_next_token_ids.append(next_token_ids) # share some inference info with the main model draft_model_input0, draft_model_input1 = model_input0, model_input1 - draft_hidden0 = model_output0.collector.spec_hidden - draft_hidden1 = model_output1.collector.spec_hidden + draft_hidden0 = model_output0.mtp_collector.spec_hidden + draft_hidden1 = model_output1.mtp_collector.spec_hidden draft_next_token_ids_gpu0 = self._build_padded_next_token_ids( token_ids=next_token_ids, @@ -974,8 +974,8 @@ def _draft_decode_vanilla_overlap( draft_model_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( draft_model_input0, draft_model_input1 ) - draft_hidden0 = draft_model_output0.collector.spec_hidden - draft_hidden1 = draft_model_output1.collector.spec_hidden + draft_hidden0 = draft_model_output0.mtp_collector.spec_hidden + draft_hidden1 = draft_model_output1.mtp_collector.spec_hidden draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index 522aec6c17..40f78681f2 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -44,7 +44,7 @@ def propose_next( ) self.extend_draft_kv_cache( main_model_input=main_model_input, - target_hidden=main_model_output.collector.spec_hidden, + target_hidden=main_model_output.mtp_collector.spec_hidden, ) draft_output = draft_model.forward(draft_input) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index bbe2da75e6..2abd4ffd02 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -54,19 +54,19 @@ def propose_next( ) self.extend_draft_kv_cache( main_model_input=main_model_input, - target_hidden=main_model_output.collector.spec_hidden, + target_hidden=main_model_output.mtp_collector.spec_hidden, ) draft_output = draft_model.forward(draft_input) - if draft_output.collector.draft_token_ids is None: + if draft_output.mtp_collector.draft_token_ids is None: flat_draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) else: - flat_draft_token_ids = draft_output.collector.draft_token_ids + flat_draft_token_ids = draft_output.mtp_collector.draft_token_ids block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) proposal_token_ids[accepted_tail_rows, 1:] = block_draft_token_ids[:, :draft_step] if self.enable_dynmaic_mtp: - confidence_logits = draft_output.collector.confidence_logits + confidence_logits = draft_output.mtp_collector.confidence_logits if confidence_logits is None: raise RuntimeError("DSpark dynamic verify requires confidence head logits") # Match the clamp used by the GPU dynamic row selector before it diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py index 82a3c5989a..cb5c9ce299 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py @@ -37,7 +37,7 @@ def build_draft_state_from_prefill( prepare_mtp_prefill_inputs( model_input=target_model_input, b_next_token_ids=next_token_ids, - mtp_draft_input_hiddens=target_model_output.collector.spec_hidden, + mtp_draft_input_hiddens=target_model_output.mtp_collector.spec_hidden, ) self.backend.draft_models[0].forward(target_model_input) @@ -55,12 +55,12 @@ def build_draft_state_from_prefill_overlap( prepare_mtp_prefill_inputs( model_input=target_model_input0, b_next_token_ids=next_token_ids0, - mtp_draft_input_hiddens=target_model_output0.collector.spec_hidden, + mtp_draft_input_hiddens=target_model_output0.mtp_collector.spec_hidden, ) prepare_mtp_prefill_inputs( model_input=target_model_input1, b_next_token_ids=next_token_ids1, - mtp_draft_input_hiddens=target_model_output1.collector.spec_hidden, + mtp_draft_input_hiddens=target_model_output1.mtp_collector.spec_hidden, ) self.backend.draft_models[0].microbatch_overlap_prefill(target_model_input0, target_model_input1) @@ -143,7 +143,7 @@ def propose_next( self.prepare_verify_extend_input( model_input=main_model_input, input_ids=next_token_ids, - target_hidden=main_model_output.collector.spec_hidden, + target_hidden=main_model_output.mtp_collector.spec_hidden, ) extend_output = draft_model.forward(main_model_input) @@ -154,7 +154,7 @@ def propose_next( else: draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) proposal_token_ids[accepted_tail_rows, 1] = draft_token_ids - draft_hidden = extend_output.collector.spec_hidden.index_select(0, accepted_tail_rows) + draft_hidden = extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows) if draft_step == 1: return SpecProposal( @@ -196,7 +196,7 @@ def propose_next( else: draft_token_ids = self._gen_argmax_token_ids(draft_output) proposal_token_ids[accepted_tail_rows, step + 1] = draft_token_ids - draft_hidden = draft_output.collector.spec_hidden + draft_hidden = draft_output.mtp_collector.spec_hidden draft_seq_lens.add_(1) return SpecProposal( @@ -262,7 +262,7 @@ def propose_next_overlap( self.prepare_verify_extend_input( model_input=model_input, input_ids=token_ids, - target_hidden=model_output.collector.spec_hidden, + target_hidden=model_output.mtp_collector.spec_hidden, ) verify_row_count = real_verify_rows0 + real_verify_rows1 @@ -288,7 +288,7 @@ def propose_next_overlap( accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) draft_token_ids_by_batch.append(draft_token_ids) - draft_hiddens_by_batch.append(extend_output.collector.spec_hidden.index_select(0, accepted_tail_rows)) + draft_hiddens_by_batch.append(extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows)) draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) proposal_rows = accepted_tail_rows[:real_request_count] + proposal_row_offsets[batch_index] @@ -346,7 +346,7 @@ def propose_next_overlap( for batch_index, draft_output in enumerate(draft_outputs): draft_token_ids = self._gen_argmax_token_ids(draft_output) draft_token_ids_by_batch[batch_index] = draft_token_ids - draft_hiddens_by_batch[batch_index] = draft_output.collector.spec_hidden + draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden draft_seq_lens_by_batch[batch_index].add_(1) real_request_count = real_request_counts[batch_index] proposal_token_ids[proposal_rows_by_batch[batch_index], step + 1] = draft_token_ids[:real_request_count] @@ -406,7 +406,10 @@ def propose_next_overlap( ) draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] - draft_hiddens_by_batch = [main_model_output0.collector.spec_hidden, main_model_output1.collector.spec_hidden] + draft_hiddens_by_batch = [ + main_model_output0.mtp_collector.spec_hidden, + main_model_output1.mtp_collector.spec_hidden, + ] draft_model = self.backend.draft_models[0] hold_mem_index = self.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX @@ -438,7 +441,7 @@ def propose_next_overlap( ).view(-1) draft_token_ids_by_batch[batch_index] = self._gen_argmax_token_ids(draft_output) - draft_hiddens_by_batch[batch_index] = draft_output.collector.spec_hidden + draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids_by_batch[0][:real_verify_rows0] proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids_by_batch[1][:real_verify_rows1] diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py index b60ab20ea7..80a7ce41be 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py @@ -40,7 +40,7 @@ def build_draft_state_from_prefill( target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - target_hidden = target_model_output.collector.spec_hidden + target_hidden = target_model_output.mtp_collector.spec_hidden if target_hidden.numel() == 0: return diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py index aa0b94f49d..a5fad08c84 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py @@ -34,7 +34,7 @@ def build_draft_state_from_prefill( ) -> None: from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - draft_hidden = target_model_output.collector.spec_hidden + draft_hidden = target_model_output.mtp_collector.spec_hidden draft_token_ids = next_token_ids for draft_model in self.backend.draft_models: prepare_mtp_prefill_inputs( @@ -43,7 +43,7 @@ def build_draft_state_from_prefill( mtp_draft_input_hiddens=draft_hidden, ) draft_output = draft_model.forward(target_model_input) - draft_hidden = draft_output.collector.spec_hidden + draft_hidden = draft_output.mtp_collector.spec_hidden draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) def build_draft_state_from_prefill_overlap( @@ -58,8 +58,8 @@ def build_draft_state_from_prefill_overlap( from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs draft_hiddens_by_batch = [ - target_model_output0.collector.spec_hidden, - target_model_output1.collector.spec_hidden, + target_model_output0.mtp_collector.spec_hidden, + target_model_output1.mtp_collector.spec_hidden, ] draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] @@ -79,7 +79,7 @@ def build_draft_state_from_prefill_overlap( target_model_input1, ) for batch_index, draft_output in enumerate(draft_outputs): - draft_hiddens_by_batch[batch_index] = draft_output.collector.spec_hidden + draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden draft_token_ids_by_batch[batch_index] = self.backend._gen_argmax_token_ids(draft_output) def propose_next( @@ -93,7 +93,7 @@ def propose_next( ) -> SpecProposal: verify_row_count = int(next_token_ids.shape[0]) draft_token_ids = next_token_ids - draft_hidden = main_model_output.collector.spec_hidden + draft_hidden = main_model_output.mtp_collector.spec_hidden proposal_token_ids = next_token_ids.new_empty((verify_row_count, draft_step + 1)) proposal_token_ids[:, 0] = next_token_ids schedule_scores = ( @@ -111,7 +111,7 @@ def propose_next( main_model_input.input_ids = draft_token_ids main_model_input.mtp_draft_input_hiddens = draft_hidden draft_output = draft_model.forward(main_model_input) - draft_hidden = draft_output.collector.spec_hidden + draft_hidden = draft_output.mtp_collector.spec_hidden if self.enable_dynmaic_mtp: draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) schedule_scores[:, step] = draft_token_probs diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index 4f09cf7099..18333d763c 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -8,24 +8,24 @@ def test_decode_unpad_slices_spec_output_with_logits(): model = TpPartBaseModel.__new__(TpPartBaseModel) output = ModelOutput( logits=torch.arange(24).view(6, 4), - collector=ModelMtpOutputCollector(spec_hidden=torch.arange(18).view(6, 3)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.arange(18).view(6, 3)), ) unpadded = model._create_unpad_decode_model_output(output, origin_batch_size=4) assert unpadded.logits.shape == (4, 4) - assert unpadded.collector.spec_hidden.shape == (4, 3) + assert unpadded.mtp_collector.spec_hidden.shape == (4, 3) # Unpadding returns a shallow output copy and leaves the graph-owned # tensors on the original ModelOutput intact. assert output.logits.shape == (6, 4) - assert output.collector.spec_hidden.shape == (6, 3) + assert output.mtp_collector.spec_hidden.shape == (6, 3) def test_prefill_unpad_uses_token_rows_for_spec_hidden(): model = TpPartBaseModel.__new__(TpPartBaseModel) output = ModelOutput( logits=torch.arange(20).view(5, 4), - collector=ModelMtpOutputCollector(spec_hidden=torch.arange(24).view(8, 3)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.arange(24).view(8, 3)), prompt_logics=torch.arange(32).view(8, 4), ) @@ -36,5 +36,5 @@ def test_prefill_unpad_uses_token_rows_for_spec_hidden(): ) assert unpadded.logits.shape == (3, 4) - assert unpadded.collector.spec_hidden.shape == (6, 3) + assert unpadded.mtp_collector.spec_hidden.shape == (6, 3) assert unpadded.prompt_logics.shape == (6, 4) diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index f967e86ffb..6e4a8f9929 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -17,7 +17,7 @@ def test_dspark_decode_unpad_uses_common_output_and_slices_mtp_fields(model_clas model = model_class.__new__(model_class) output = ModelOutput( logits=torch.arange(48).view(12, 4), - collector=ModelMtpOutputCollector( + mtp_collector=ModelMtpOutputCollector( spec_hidden=torch.arange(36).view(12, 3), confidence_logits=torch.arange(12).view(3, 4), draft_token_ids=torch.arange(12), @@ -28,19 +28,19 @@ def test_dspark_decode_unpad_uses_common_output_and_slices_mtp_fields(model_clas assert isinstance(unpadded, ModelOutput) assert unpadded.logits.shape == (8, 4) - assert unpadded.collector.spec_hidden.shape == (8, 3) - assert unpadded.collector.confidence_logits.shape == (2, 4) - assert unpadded.collector.draft_token_ids.shape == (8,) + assert unpadded.mtp_collector.spec_hidden.shape == (8, 3) + assert unpadded.mtp_collector.confidence_logits.shape == (2, 4) + assert unpadded.mtp_collector.draft_token_ids.shape == (8,) assert output.logits.shape == (12, 4) - assert output.collector.confidence_logits.shape == (3, 4) - assert output.collector.draft_token_ids.shape == (12,) + assert output.mtp_collector.confidence_logits.shape == (3, 4) + assert output.mtp_collector.draft_token_ids.shape == (12,) def test_common_output_no_ref_conversion_includes_dspark_fields(monkeypatch): monkeypatch.setattr(batch_objs, "tensor_to_no_ref_tensor", torch.clone) output = ModelOutput( logits=torch.ones((2, 4)), - collector=ModelMtpOutputCollector( + mtp_collector=ModelMtpOutputCollector( spec_hidden=torch.ones((2, 3)), confidence_logits=torch.ones((1, 2)), draft_token_ids=torch.ones((2,), dtype=torch.int64), @@ -48,18 +48,18 @@ def test_common_output_no_ref_conversion_includes_dspark_fields(monkeypatch): ) original_ptrs = ( output.logits.data_ptr(), - output.collector.spec_hidden.data_ptr(), - output.collector.confidence_logits.data_ptr(), - output.collector.draft_token_ids.data_ptr(), + output.mtp_collector.spec_hidden.data_ptr(), + output.mtp_collector.confidence_logits.data_ptr(), + output.mtp_collector.draft_token_ids.data_ptr(), ) output.to_no_ref_tensor() converted_ptrs = ( output.logits.data_ptr(), - output.collector.spec_hidden.data_ptr(), - output.collector.confidence_logits.data_ptr(), - output.collector.draft_token_ids.data_ptr(), + output.mtp_collector.spec_hidden.data_ptr(), + output.mtp_collector.confidence_logits.data_ptr(), + output.mtp_collector.draft_token_ids.data_ptr(), ) assert all(converted != original for converted, original in zip(converted_ptrs, original_ptrs)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index bb7846ab45..a0c663a08c 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -23,7 +23,7 @@ def microbatch_overlap_prefill(self, input0, input1): return tuple( ModelOutput( logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), - collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), ) for model_input in (input0, input1) ) @@ -34,7 +34,7 @@ def microbatch_overlap_decode(self, input0, input1): return tuple( ModelOutput( logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), - collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), ) for model_input in (input0, input1) ) @@ -79,14 +79,14 @@ def test_overlap_eagle_keeps_fixed_verify_layout(): proposal = proposer.propose_next_overlap( main_model_input0=model_input0, main_model_output0=ModelOutput( - logits=torch.empty((6, 1)), collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), next_token_ids0=torch.arange(6, dtype=torch.int64), real_verify_rows0=3, accept_len0=torch.tensor([2, 1], dtype=torch.int32), main_model_input1=model_input1, main_model_output1=ModelOutput( - logits=torch.empty((6, 1)), collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), next_token_ids1=torch.arange(10, 16, dtype=torch.int64), real_verify_rows1=6, @@ -126,14 +126,14 @@ def test_autoregressive_eagle_reuses_overlap_inputs(): proposal = proposer.propose_next_overlap( main_model_input0=model_input0, main_model_output0=ModelOutput( - logits=torch.empty((6, 1)), collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), next_token_ids0=torch.arange(6, dtype=torch.int64), real_verify_rows0=3, accept_len0=torch.tensor([2, 1], dtype=torch.int32), main_model_input1=model_input1, main_model_output1=ModelOutput( - logits=torch.empty((6, 1)), collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), next_token_ids1=torch.arange(10, 16, dtype=torch.int64), real_verify_rows1=6, From 3dce9acd4a8c7647cf2bf1641e9f369ac32796ed Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 06:46:51 +0000 Subject: [PATCH 037/103] refactor: clarify proposal coverage state --- .../router/model_infer/mtp_speculative/planner.py | 12 ++++++------ .../model_infer/mtp_speculative/test_planner.py | 12 ++++++------ 2 files changed, 12 insertions(+), 12 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner.py b/lightllm/server/router/model_infer/mtp_speculative/planner.py index 41e18106da..854a81bbac 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner.py @@ -32,7 +32,7 @@ class SpecDecodePlan: # False when the current batch contains requests without a proposal from # the previous iteration. Such an iteration does not represent one # well-defined LightSpec runtime configuration. - record_progress: bool = True + all_reqs_have_proposals: bool = True @property def is_dynamic(self) -> bool: @@ -121,7 +121,7 @@ def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> # This bounds Verify by the candidate pool that physically exists. available_batch_size = req_num + proposal_req_num * pre_draft_step max_batch_size = min(original_batch_size, available_batch_size) - record_progress = proposal_req_num == req_num + all_reqs_have_proposals = proposal_req_num == req_num if ( not self.target_infer_costs.has_data() @@ -136,7 +136,7 @@ def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> dynamic_batch_size=max_batch_size, draft_step=self.max_draft_step, pre_draft_step=pre_draft_step, - record_progress=record_progress, + all_reqs_have_proposals=all_reqs_have_proposals, ) min_batch_size = req_num @@ -157,7 +157,7 @@ def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> dynamic_batch_size=dynamic_batch_size, draft_step=draft_step, pre_draft_step=pre_draft_step, - record_progress=record_progress, + all_reqs_have_proposals=all_reqs_have_proposals, ) def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: @@ -172,12 +172,12 @@ def update_feedback( schedule_scores=None, ) -> None: # The progress EMA records one complete-batch sample for a single - # (N, B, d) configuration. record_progress is false if any request is + # (N, B, d) configuration. all_reqs_have_proposals is false if any request is # on its first decode and therefore has no preceding proposal; its # structural accept_len=1 would otherwise bias that configuration's # progress downward. Skip the mixed batch instead of partially # updating the batch-level statistic with different N/B semantics. - if not plan.record_progress: + if not plan.all_reqs_have_proposals: return self.update_verified_batch( accept_lengths=accept_lengths, diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index f440609e97..943628377e 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -262,9 +262,9 @@ def test_lightspec_bounds_verify_to_existing_proposals(): assert cold_start.dynamic_batch_size == 2 assert mixed_batch.dynamic_batch_size == 5 assert ready_batch.dynamic_batch_size == 8 - assert not cold_start.record_progress - assert not mixed_batch.record_progress - assert ready_batch.record_progress + assert not cold_start.all_reqs_have_proposals + assert not mixed_batch.all_reqs_have_proposals + assert ready_batch.all_reqs_have_proposals def test_engine_counts_requests_with_a_previous_proposal(): @@ -282,9 +282,9 @@ def test_engine_counts_requests_with_a_previous_proposal(): ) assert mixed_plan.dynamic_batch_size == 5 - assert not mixed_plan.record_progress + assert not mixed_plan.all_reqs_have_proposals assert ready_plan.dynamic_batch_size == 8 - assert ready_plan.record_progress + assert ready_plan.all_reqs_have_proposals def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): @@ -459,7 +459,7 @@ def test_engine_skips_feedback_for_a_mixed_proposal_batch(): dynamic_batch_size=5, draft_step=3, pre_draft_step=3, - record_progress=False, + all_reqs_have_proposals=False, ) engine.update_planner_feedback( From ce151423f22251271f293b752a3e2de8752449c5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 06:55:01 +0000 Subject: [PATCH 038/103] refactor: split speculative planners into package --- .../model_infer/mtp_speculative/planner.py | 510 ------------------ .../mtp_speculative/planner/__init__.py | 12 + .../mtp_speculative/planner/base.py | 103 ++++ .../mtp_speculative/planner/dspark.py | 168 ++++++ .../mtp_speculative/planner/fixed.py | 17 + .../mtp_speculative/planner/lightspec.py | 237 ++++++++ .../mtp_speculative/test_planner.py | 2 +- 7 files changed, 538 insertions(+), 511 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/planner.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/planner/base.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner.py b/lightllm/server/router/model_infer/mtp_speculative/planner.py deleted file mode 100644 index 854a81bbac..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/planner.py +++ /dev/null @@ -1,510 +0,0 @@ -from __future__ import annotations - -from collections import deque -from dataclasses import dataclass -from typing import TYPE_CHECKING, Dict, List, Optional, Tuple - -import numpy as np -from sortedcontainers import SortedDict - -if TYPE_CHECKING: - from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer - - -@dataclass(frozen=True) -class SpecDecodePlan: - """Planner decision for one target decode iteration. - - Fixed scheduling uses the full speculative-expanded target batch: - - dynamic_batch_size is None - - draft_step == max_draft_step - - Dynamic speculative scheduling may compact target rows before forward: - - dynamic_batch_size is the selected target row count - - draft_step is the candidate length to generate after target verify - - pre_draft_step describes the previous iteration and controls whether - GPU verify sync can be skipped - """ - - dynamic_batch_size: Optional[int] - draft_step: int - pre_draft_step: int - # False when the current batch contains requests without a proposal from - # the previous iteration. Such an iteration does not represent one - # well-defined LightSpec runtime configuration. - all_reqs_have_proposals: bool = True - - @property - def is_dynamic(self) -> bool: - return self.dynamic_batch_size is not None - - @property - def skip_verify_sync(self) -> bool: - return self.is_dynamic and self.pre_draft_step == 0 - - def filter_reqs(self, reqs: List, selected_row_mask_cpu) -> List: - return [req for req, selected in zip(reqs, selected_row_mask_cpu.tolist()) if selected] - - -class FixedSpecPlanner: - """Planner for fixed-width speculative decoding.""" - - def __init__(self, max_draft_step: int) -> None: - self.max_draft_step = int(max_draft_step) - - def plan(self, req_num: int | None = None, original_batch_size: int | None = None) -> SpecDecodePlan: - return SpecDecodePlan( - dynamic_batch_size=None, - draft_step=self.max_draft_step, - pre_draft_step=self.max_draft_step, - ) - - -class LightSpecPlanner: - """Choose the current verify budget and the next draft configuration. - - LightSpec evaluates a runtime configuration as ``(N, B, d)``: logical - requests, physical target rows, and draft configuration. Planning follows - the paper's budget-then-fill order. It chooses the current ``B`` within the - proposal produced by ``pre_draft_step``. The next ``draft_step`` is selected - from self-consistent ``(B, d)`` configurations for the following iteration, - so a short current proposal cannot prevent the planner from drafting deeper - again. Candidate identities belong to the GPU Fill stage. - - Each configuration minimizes estimated milliseconds per committed token: - ``(target_cost(B) + draft_cost(N, B, d)) / expected_progress(N, B, d)``. - Progress is learned from one normalized batch observation ``sum(accept_len) - / B`` per iteration, then smoothed with an EMA. It is never updated once per - request; doing so would make adaptation depend on concurrency. - - The proposer supplies its valid draft configurations and complete draft - cost. The planner therefore stays independent of the proposal algorithm and - its physical execution layout. - """ - - def __init__( - self, - max_draft_step: int, - draft_cost_provider: "BaseSpecProposer", - ) -> None: - self.max_draft_step = int(max_draft_step) - self.draft_cost_provider = draft_cost_provider - self.draft_steps = draft_cost_provider.get_draft_steps() - - self.target_infer_costs = _InferCostMsTable() - self.draft_infer_costs = _InferCostMsTable() - - # Each observation is the normalized committed progress U / B from one - # complete batch. The draft configuration is part of the key because a - # deeper candidate pool may improve Fill even at the same (N, B). - self.progress_ema_by_config: Dict[Tuple[int, int, int], _EMAValue] = {} - # Full-width verification provides an unbiased survival probability - # for every draft depth. It is the fallback for configurations that - # have not produced their own progress observation yet. - self.prefix_survival_by_depth: List[Optional[_EMAValue]] = [None] * self.max_draft_step - - # The current verify width is bounded by the proposal built last time. - self.pre_draft_step = self.max_draft_step - - def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> SpecDecodePlan: - pre_draft_step = self.pre_draft_step - if req_num == 0: - self.pre_draft_step = self.max_draft_step - return SpecDecodePlan( - dynamic_batch_size=0, - draft_step=self.max_draft_step, - pre_draft_step=pre_draft_step, - ) - - # A request entering its first decode has only the guaranteed target - # row. Existing requests additionally expose pre_draft_step candidates. - # This bounds Verify by the candidate pool that physically exists. - available_batch_size = req_num + proposal_req_num * pre_draft_step - max_batch_size = min(original_batch_size, available_batch_size) - all_reqs_have_proposals = proposal_req_num == req_num - - if ( - not self.target_infer_costs.has_data() - or not self.draft_infer_costs.has_data() - or not self.progress_ema_by_config - ): - # Costs come from CUDA Graph capture, while progress requires one - # real verification batch. Keep the available proposal intact until - # both parts of J(N, B, d) have an observation. - self.pre_draft_step = self.max_draft_step - return SpecDecodePlan( - dynamic_batch_size=max_batch_size, - draft_step=self.max_draft_step, - pre_draft_step=pre_draft_step, - all_reqs_have_proposals=all_reqs_have_proposals, - ) - - min_batch_size = req_num - batch_sizes = self.target_infer_costs.get_batch_size_keys_between(min_batch_size, max_batch_size) - costs = [ - self._get_cost_ms( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=pre_draft_step, - ) - for dynamic_batch_size in batch_sizes - ] - dynamic_batch_size = batch_sizes[np.argmin(costs)] - draft_step = self._select_draft_step(req_num=req_num) - - self.pre_draft_step = draft_step - return SpecDecodePlan( - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - pre_draft_step=pre_draft_step, - all_reqs_have_proposals=all_reqs_have_proposals, - ) - - def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: - cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs - cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - - def update_feedback( - self, - plan: SpecDecodePlan, - req_num: int, - accept_lengths, - schedule_scores=None, - ) -> None: - # The progress EMA records one complete-batch sample for a single - # (N, B, d) configuration. all_reqs_have_proposals is false if any request is - # on its first decode and therefore has no preceding proposal; its - # structural accept_len=1 would otherwise bias that configuration's - # progress downward. Skip the mixed batch instead of partially - # updating the batch-level statistic with different N/B semantics. - if not plan.all_reqs_have_proposals: - return - self.update_verified_batch( - accept_lengths=accept_lengths, - req_num=req_num, - dynamic_batch_size=plan.dynamic_batch_size, - verified_draft_step=plan.pre_draft_step, - ) - - def update_verified_batch( - self, - accept_lengths, - req_num: int, - dynamic_batch_size: int, - verified_draft_step: int, - ) -> None: - """Record one batch-level progress sample for the verified configuration.""" - - accept_lengths = np.asarray(accept_lengths) - if accept_lengths.size == 0: - return - - config = (int(req_num), int(dynamic_batch_size), int(verified_draft_step)) - progress = float(accept_lengths.sum()) / dynamic_batch_size - if config not in self.progress_ema_by_config: - self.progress_ema_by_config[config] = _EMAValue( - decay=0.9, - init_value=progress, - ) - self.progress_ema_by_config[config].update(progress) - - if dynamic_batch_size == req_num * (verified_draft_step + 1): - for depth in range(1, verified_draft_step + 1): - survival = float(np.mean(accept_lengths > depth)) - survival_ema = self.prefix_survival_by_depth[depth - 1] - if survival_ema is None: - survival_ema = _EMAValue(decay=0.9, init_value=survival) - self.prefix_survival_by_depth[depth - 1] = survival_ema - survival_ema.update(survival) - - def _select_draft_step(self, req_num: int) -> int: - """Choose the best self-consistent next proposal depth. - - The current target batch is bounded by ``pre_draft_step``, but the - proposal generated now is consumed by the next iteration. Evaluate - each candidate depth with the verify widths that depth can create so a - short current proposal cannot permanently prevent deeper drafting. - """ - - best_cost_ms = float("inf") - best_draft_step = self.draft_steps[0] - for draft_step in self.draft_steps: - max_batch_size = req_num * (draft_step + 1) - batch_sizes = self.target_infer_costs.get_batch_size_keys_between(req_num, max_batch_size) - for dynamic_batch_size in batch_sizes: - cost_ms = self._get_cost_ms( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - ) - if cost_ms < best_cost_ms: - best_cost_ms = cost_ms - best_draft_step = draft_step - return best_draft_step - - def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: - """Estimate milliseconds per committed token for one configuration.""" - - accept_ratio = self._estimate_progress( - req_num=req_num, - dynamic_batch_size=dynamic_batch_size, - draft_step=draft_step, - ) - total_time = self.target_infer_costs.get(dynamic_batch_size) + self.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=self.draft_infer_costs, - req_num=req_num, - verify_batch_size=dynamic_batch_size, - draft_step=draft_step, - ) - token_num = min((dynamic_batch_size * accept_ratio), req_num * (draft_step + 1)) - token_num = max(token_num, req_num) - cost_ms = total_time / token_num - return cost_ms - - def _estimate_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: - config = (int(req_num), int(dynamic_batch_size), int(draft_step)) - ema = self.progress_ema_by_config.get(config) - return ema.get() if ema is not None else self._estimate_prefix_progress(*config) - - def _estimate_prefix_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: - if not any(self.prefix_survival_by_depth): - return 1.0 - - remaining_draft_rows = dynamic_batch_size - req_num - expected_tokens = float(req_num) - for depth in range(min(draft_step, self.max_draft_step)): - selected_rows = min(req_num, remaining_draft_rows) - if selected_rows <= 0: - break - survival_ema = self.prefix_survival_by_depth[depth] - survival = 1.0 if survival_ema is None else survival_ema.get() - expected_tokens += selected_rows * survival - remaining_draft_rows -= selected_rows - return expected_tokens / dynamic_batch_size - - -class DSparkPlanner: - """DSpark's confidence-based verify-capacity planner. - - DSpark always drafts a complete block. Confidence from one proposal selects - the target verify capacity two iterations later. - """ - - def __init__(self, max_draft_step: int, draft_cost_provider: "BaseSpecProposer") -> None: - self.max_draft_step = int(max_draft_step) - self.draft_cost_provider = draft_cost_provider - self.target_infer_costs = _InferCostMsTable() - self.draft_infer_costs = _InferCostMsTable() - self._pending_verify_batch_sizes = deque(maxlen=2) - - def plan(self, req_num: int, original_batch_size: int) -> SpecDecodePlan: - if req_num == 0: - return SpecDecodePlan( - dynamic_batch_size=0, - draft_step=self.max_draft_step, - pre_draft_step=self.max_draft_step, - ) - - full_batch_size = original_batch_size - dynamic_batch_size = full_batch_size - if self.target_infer_costs.has_data() and self.draft_infer_costs.has_data(): - delayed_batch_size = self._pop_delayed_batch_size( - req_num=req_num, - max_batch_size=full_batch_size, - ) - if delayed_batch_size is not None: - dynamic_batch_size = delayed_batch_size - - return SpecDecodePlan( - dynamic_batch_size=dynamic_batch_size, - draft_step=self.max_draft_step, - pre_draft_step=self.max_draft_step, - ) - - def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: - cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs - cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - - def update_feedback( - self, - plan: SpecDecodePlan, - req_num: int, - accept_lengths, - schedule_scores=None, - ) -> None: - if schedule_scores is not None: - self.update_confidence_probs( - confidence_probs=schedule_scores, - req_num=req_num, - ) - - def update_confidence_probs(self, confidence_probs, req_num: int) -> None: - """Record a confidence-derived future capacity estimate. - - The current decode step still routes rows by the current probabilities - stored in req_to_next_token_scores. This queue is only used to choose - the future capacity K after a two-step delay, matching the asynchronous - scheduler constraint described by DSpark. - """ - - if req_num <= 0: - return - if not self.target_infer_costs.has_data() or not self.draft_infer_costs.has_data(): - return - - probs = np.asarray(confidence_probs, dtype=np.float64) - if probs.ndim != 2 or probs.shape[1] == 0: - return - - draft_confidence_probs = probs[:, : self.max_draft_step] - if draft_confidence_probs.size == 0: - return - - # Confidence is scattered onto one accepted-tail row per request; - # unused verify rows remain zero. - valid_rows = np.any(draft_confidence_probs > 0.0, axis=1) - if not np.any(valid_rows): - return - - conditional_probs = np.clip(draft_confidence_probs[valid_rows], 0.01, 0.99) - survival_scores = np.cumprod(conditional_probs, axis=1) - dynamic_batch_size = self._select_dynamic_batch_size_from_survival_scores( - req_num=int(req_num), - survival_scores=survival_scores, - ) - self._pending_verify_batch_sizes.append(dynamic_batch_size) - - def _pop_delayed_batch_size(self, req_num: int, max_batch_size: int) -> Optional[int]: - if len(self._pending_verify_batch_sizes) < 2: - return None - predicted_batch_size = int(self._pending_verify_batch_sizes.popleft()) - return min(max(predicted_batch_size, int(req_num)), int(max_batch_size)) - - def _select_dynamic_batch_size_from_survival_scores( - self, - req_num: int, - survival_scores: np.ndarray, - ) -> int: - flat_survival_scores = survival_scores.reshape(-1) - max_batch_size = int(req_num + flat_survival_scores.shape[0]) - if flat_survival_scores.shape[0] == 0: - return int(req_num) - - candidate_batch_sizes = set(self.target_infer_costs.get_batch_size_keys_between(req_num, max_batch_size)) - candidate_batch_sizes.add(int(req_num)) - candidate_batch_sizes.add(max_batch_size) - - candidate_batch_sizes = { - min(max(int(dynamic_batch_size), int(req_num)), max_batch_size) - for dynamic_batch_size in candidate_batch_sizes - } - selected_draft_counts = sorted( - {int(dynamic_batch_size) - int(req_num) for dynamic_batch_size in candidate_batch_sizes} - ) - expected_accepts_by_count = self._topk_prefix_sums( - values=flat_survival_scores, - counts=selected_draft_counts, - ) - - best_batch_size = int(req_num) - best_throughput = -float("inf") - for dynamic_batch_size in sorted(candidate_batch_sizes): - selected_draft_count = dynamic_batch_size - int(req_num) - expected_tokens = float(req_num) + float(expected_accepts_by_count[selected_draft_count]) - round_ms = self.target_infer_costs.get(dynamic_batch_size) + self.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=self.draft_infer_costs, - req_num=req_num, - verify_batch_size=dynamic_batch_size, - draft_step=self.max_draft_step, - ) - throughput = expected_tokens / max(round_ms, 1e-6) - if throughput > best_throughput: - best_throughput = throughput - best_batch_size = dynamic_batch_size - return best_batch_size - - @staticmethod - def _topk_prefix_sums(values: np.ndarray, counts: List[int]) -> Dict[int, float]: - """Return sum(top-k(values)) only for the requested k values.""" - - if not counts: - return {} - - flat_values = np.asarray(values, dtype=np.float64).reshape(-1) - value_count = int(flat_values.shape[0]) - normalized_counts = sorted({min(max(int(count), 0), value_count) for count in counts}) - if value_count == 0: - return {count: 0.0 for count in normalized_counts} - - prefix_sums = np.concatenate(([0.0], np.cumsum(np.sort(flat_values)[::-1], dtype=np.float64))) - return {count: float(prefix_sums[count]) for count in normalized_counts} - - -class _InferCostMsTable: - def __init__(self) -> None: - self.infer_cost_ms_table = SortedDict() - - def update(self, batch_size: int, infer_cost_ms: float) -> None: - self.infer_cost_ms_table[int(batch_size)] = float(infer_cost_ms) - - def has_data(self) -> bool: - return len(self.infer_cost_ms_table) > 0 - - def get(self, batch_size: int) -> float: - batch_size = int(batch_size) - - if len(self.infer_cost_ms_table) == 0: - return batch_size * 1000.0 - - if batch_size in self.infer_cost_ms_table: - return self.infer_cost_ms_table[batch_size] - - max_batch_size = self.infer_cost_ms_table.peekitem(-1)[0] - if batch_size > max_batch_size: - max_infer_cost_ms = self.infer_cost_ms_table.peekitem(-1)[1] - # 超过最大 graph 范围时使用高惩罚,使调度器倾向于关闭 speculative draft。 - return max_infer_cost_ms + (batch_size - max_batch_size) * 1000.0 - - index = self.infer_cost_ms_table.bisect_left(batch_size) - return self.infer_cost_ms_table.peekitem(index)[1] - - def estimate(self, batch_size: int) -> float: - """Estimate an uncaptured batch without applying the graph-miss penalty.""" - - batch_size = int(batch_size) - max_batch_size, max_cost_ms = self.infer_cost_ms_table.peekitem(-1) - if batch_size <= max_batch_size: - return self.get(batch_size) - return max_cost_ms * batch_size / max_batch_size - - def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> List[int]: - start = min(int(batch_size1), int(batch_size2)) - end = max(int(batch_size1), int(batch_size2)) - batch_sizes = set(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) - batch_sizes.update((start, end)) - return sorted(batch_sizes) - - -class _EMAValue: - def __init__(self, decay: float, init_value: float) -> None: - self.decay = decay - self.value = init_value - self.update_count = 0 - - def update(self, new_value: float): - self.update_count += 1 - self.value = self.decay * self.value + (1.0 - self.decay) * new_value - - def get(self) -> float: - return self.value - - def get_count(self) -> int: - return self.update_count - - -__all__ = [ - "DSparkPlanner", - "FixedSpecPlanner", - "LightSpecPlanner", - "SpecDecodePlan", -] diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py new file mode 100644 index 0000000000..a48dcb7444 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py @@ -0,0 +1,12 @@ +from lightllm.server.router.model_infer.mtp_speculative.planner.base import SpecDecodePlan +from lightllm.server.router.model_infer.mtp_speculative.planner.dspark import DSparkPlanner +from lightllm.server.router.model_infer.mtp_speculative.planner.fixed import FixedSpecPlanner +from lightllm.server.router.model_infer.mtp_speculative.planner.lightspec import LightSpecPlanner + + +__all__ = [ + "DSparkPlanner", + "FixedSpecPlanner", + "LightSpecPlanner", + "SpecDecodePlan", +] diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py new file mode 100644 index 0000000000..f6e1304de8 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import List, Optional + +from sortedcontainers import SortedDict + + +@dataclass(frozen=True) +class SpecDecodePlan: + """Planner decision for one target decode iteration. + + Fixed scheduling uses the full speculative-expanded target batch: + - dynamic_batch_size is None + - draft_step == max_draft_step + + Dynamic speculative scheduling may compact target rows before forward: + - dynamic_batch_size is the selected target row count + - draft_step is the candidate length to generate after target verify + - pre_draft_step describes the previous iteration and controls whether + GPU verify sync can be skipped + """ + + dynamic_batch_size: Optional[int] + draft_step: int + pre_draft_step: int + # False when the current batch contains requests without a proposal from + # the previous iteration. Such an iteration does not represent one + # well-defined LightSpec runtime configuration. + all_reqs_have_proposals: bool = True + + @property + def is_dynamic(self) -> bool: + return self.dynamic_batch_size is not None + + @property + def skip_verify_sync(self) -> bool: + return self.is_dynamic and self.pre_draft_step == 0 + + def filter_reqs(self, reqs: List, selected_row_mask_cpu) -> List: + return [req for req, selected in zip(reqs, selected_row_mask_cpu.tolist()) if selected] + + +class _InferCostMsTable: + def __init__(self) -> None: + self.infer_cost_ms_table = SortedDict() + + def update(self, batch_size: int, infer_cost_ms: float) -> None: + self.infer_cost_ms_table[int(batch_size)] = float(infer_cost_ms) + + def has_data(self) -> bool: + return len(self.infer_cost_ms_table) > 0 + + def get(self, batch_size: int) -> float: + batch_size = int(batch_size) + + if len(self.infer_cost_ms_table) == 0: + return batch_size * 1000.0 + + if batch_size in self.infer_cost_ms_table: + return self.infer_cost_ms_table[batch_size] + + max_batch_size = self.infer_cost_ms_table.peekitem(-1)[0] + if batch_size > max_batch_size: + max_infer_cost_ms = self.infer_cost_ms_table.peekitem(-1)[1] + # 超过最大 graph 范围时使用高惩罚,使调度器倾向于关闭 speculative draft。 + return max_infer_cost_ms + (batch_size - max_batch_size) * 1000.0 + + index = self.infer_cost_ms_table.bisect_left(batch_size) + return self.infer_cost_ms_table.peekitem(index)[1] + + def estimate(self, batch_size: int) -> float: + """Estimate an uncaptured batch without applying the graph-miss penalty.""" + + batch_size = int(batch_size) + max_batch_size, max_cost_ms = self.infer_cost_ms_table.peekitem(-1) + if batch_size <= max_batch_size: + return self.get(batch_size) + return max_cost_ms * batch_size / max_batch_size + + def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> List[int]: + start = min(int(batch_size1), int(batch_size2)) + end = max(int(batch_size1), int(batch_size2)) + batch_sizes = set(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) + batch_sizes.update((start, end)) + return sorted(batch_sizes) + + +class _EMAValue: + def __init__(self, decay: float, init_value: float) -> None: + self.decay = decay + self.value = init_value + self.update_count = 0 + + def update(self, new_value: float): + self.update_count += 1 + self.value = self.decay * self.value + (1.0 - self.decay) * new_value + + def get(self) -> float: + return self.value + + def get_count(self) -> int: + return self.update_count diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py new file mode 100644 index 0000000000..b68211d709 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -0,0 +1,168 @@ +from __future__ import annotations + +from collections import deque +from typing import TYPE_CHECKING, Dict, List, Optional + +import numpy as np + +from lightllm.server.router.model_infer.mtp_speculative.planner.base import SpecDecodePlan, _InferCostMsTable + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer + + +class DSparkPlanner: + """DSpark's confidence-based verify-capacity planner. + + DSpark always drafts a complete block. Confidence from one proposal selects + the target verify capacity two iterations later. + """ + + def __init__(self, max_draft_step: int, draft_cost_provider: "BaseSpecProposer") -> None: + self.max_draft_step = int(max_draft_step) + self.draft_cost_provider = draft_cost_provider + self.target_infer_costs = _InferCostMsTable() + self.draft_infer_costs = _InferCostMsTable() + self._pending_verify_batch_sizes = deque(maxlen=2) + + def plan(self, req_num: int, original_batch_size: int) -> SpecDecodePlan: + if req_num == 0: + return SpecDecodePlan( + dynamic_batch_size=0, + draft_step=self.max_draft_step, + pre_draft_step=self.max_draft_step, + ) + + full_batch_size = original_batch_size + dynamic_batch_size = full_batch_size + if self.target_infer_costs.has_data() and self.draft_infer_costs.has_data(): + delayed_batch_size = self._pop_delayed_batch_size( + req_num=req_num, + max_batch_size=full_batch_size, + ) + if delayed_batch_size is not None: + dynamic_batch_size = delayed_batch_size + + return SpecDecodePlan( + dynamic_batch_size=dynamic_batch_size, + draft_step=self.max_draft_step, + pre_draft_step=self.max_draft_step, + ) + + def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: + cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs + cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + + def update_feedback( + self, + plan: SpecDecodePlan, + req_num: int, + accept_lengths, + schedule_scores=None, + ) -> None: + if schedule_scores is not None: + self.update_confidence_probs( + confidence_probs=schedule_scores, + req_num=req_num, + ) + + def update_confidence_probs(self, confidence_probs, req_num: int) -> None: + """Record a confidence-derived future capacity estimate. + + The current decode step still routes rows by the current probabilities + stored in req_to_next_token_scores. This queue is only used to choose + the future capacity K after a two-step delay, matching the asynchronous + scheduler constraint described by DSpark. + """ + + if req_num <= 0: + return + if not self.target_infer_costs.has_data() or not self.draft_infer_costs.has_data(): + return + + probs = np.asarray(confidence_probs, dtype=np.float64) + if probs.ndim != 2 or probs.shape[1] == 0: + return + + draft_confidence_probs = probs[:, : self.max_draft_step] + if draft_confidence_probs.size == 0: + return + + # Confidence is scattered onto one accepted-tail row per request; + # unused verify rows remain zero. + valid_rows = np.any(draft_confidence_probs > 0.0, axis=1) + if not np.any(valid_rows): + return + + conditional_probs = np.clip(draft_confidence_probs[valid_rows], 0.01, 0.99) + survival_scores = np.cumprod(conditional_probs, axis=1) + dynamic_batch_size = self._select_dynamic_batch_size_from_survival_scores( + req_num=int(req_num), + survival_scores=survival_scores, + ) + self._pending_verify_batch_sizes.append(dynamic_batch_size) + + def _pop_delayed_batch_size(self, req_num: int, max_batch_size: int) -> Optional[int]: + if len(self._pending_verify_batch_sizes) < 2: + return None + predicted_batch_size = int(self._pending_verify_batch_sizes.popleft()) + return min(max(predicted_batch_size, int(req_num)), int(max_batch_size)) + + def _select_dynamic_batch_size_from_survival_scores( + self, + req_num: int, + survival_scores: np.ndarray, + ) -> int: + flat_survival_scores = survival_scores.reshape(-1) + max_batch_size = int(req_num + flat_survival_scores.shape[0]) + if flat_survival_scores.shape[0] == 0: + return int(req_num) + + candidate_batch_sizes = set(self.target_infer_costs.get_batch_size_keys_between(req_num, max_batch_size)) + candidate_batch_sizes.add(int(req_num)) + candidate_batch_sizes.add(max_batch_size) + + candidate_batch_sizes = { + min(max(int(dynamic_batch_size), int(req_num)), max_batch_size) + for dynamic_batch_size in candidate_batch_sizes + } + selected_draft_counts = sorted( + {int(dynamic_batch_size) - int(req_num) for dynamic_batch_size in candidate_batch_sizes} + ) + expected_accepts_by_count = self._topk_prefix_sums( + values=flat_survival_scores, + counts=selected_draft_counts, + ) + + best_batch_size = int(req_num) + best_throughput = -float("inf") + for dynamic_batch_size in sorted(candidate_batch_sizes): + selected_draft_count = dynamic_batch_size - int(req_num) + expected_tokens = float(req_num) + float(expected_accepts_by_count[selected_draft_count]) + round_ms = self.target_infer_costs.get(dynamic_batch_size) + self.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=self.draft_infer_costs, + req_num=req_num, + verify_batch_size=dynamic_batch_size, + draft_step=self.max_draft_step, + ) + throughput = expected_tokens / max(round_ms, 1e-6) + if throughput > best_throughput: + best_throughput = throughput + best_batch_size = dynamic_batch_size + return best_batch_size + + @staticmethod + def _topk_prefix_sums(values: np.ndarray, counts: List[int]) -> Dict[int, float]: + """Return sum(top-k(values)) only for the requested k values.""" + + if not counts: + return {} + + flat_values = np.asarray(values, dtype=np.float64).reshape(-1) + value_count = int(flat_values.shape[0]) + normalized_counts = sorted({min(max(int(count), 0), value_count) for count in counts}) + if value_count == 0: + return {count: 0.0 for count in normalized_counts} + + prefix_sums = np.concatenate(([0.0], np.cumsum(np.sort(flat_values)[::-1], dtype=np.float64))) + return {count: float(prefix_sums[count]) for count in normalized_counts} diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py new file mode 100644 index 0000000000..177347638e --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +from lightllm.server.router.model_infer.mtp_speculative.planner.base import SpecDecodePlan + + +class FixedSpecPlanner: + """Planner for fixed-width speculative decoding.""" + + def __init__(self, max_draft_step: int) -> None: + self.max_draft_step = int(max_draft_step) + + def plan(self, req_num: int | None = None, original_batch_size: int | None = None) -> SpecDecodePlan: + return SpecDecodePlan( + dynamic_batch_size=None, + draft_step=self.max_draft_step, + pre_draft_step=self.max_draft_step, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py new file mode 100644 index 0000000000..288d76a0b2 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -0,0 +1,237 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Dict, List, Optional, Tuple + +import numpy as np + +from lightllm.server.router.model_infer.mtp_speculative.planner.base import ( + SpecDecodePlan, + _EMAValue, + _InferCostMsTable, +) + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer + + +class LightSpecPlanner: + """Choose the current verify budget and the next draft configuration. + + LightSpec evaluates a runtime configuration as ``(N, B, d)``: logical + requests, physical target rows, and draft configuration. Planning follows + the paper's budget-then-fill order. It chooses the current ``B`` within the + proposal produced by ``pre_draft_step``. The next ``draft_step`` is selected + from self-consistent ``(B, d)`` configurations for the following iteration, + so a short current proposal cannot prevent the planner from drafting deeper + again. Candidate identities belong to the GPU Fill stage. + + Each configuration minimizes estimated milliseconds per committed token: + ``(target_cost(B) + draft_cost(N, B, d)) / expected_progress(N, B, d)``. + Progress is learned from one normalized batch observation ``sum(accept_len) + / B`` per iteration, then smoothed with an EMA. It is never updated once per + request; doing so would make adaptation depend on concurrency. + + The proposer supplies its valid draft configurations and complete draft + cost. The planner therefore stays independent of the proposal algorithm and + its physical execution layout. + """ + + def __init__( + self, + max_draft_step: int, + draft_cost_provider: "BaseSpecProposer", + ) -> None: + self.max_draft_step = int(max_draft_step) + self.draft_cost_provider = draft_cost_provider + self.draft_steps = draft_cost_provider.get_draft_steps() + + self.target_infer_costs = _InferCostMsTable() + self.draft_infer_costs = _InferCostMsTable() + + # Each observation is the normalized committed progress U / B from one + # complete batch. The draft configuration is part of the key because a + # deeper candidate pool may improve Fill even at the same (N, B). + self.progress_ema_by_config: Dict[Tuple[int, int, int], _EMAValue] = {} + # Full-width verification provides an unbiased survival probability + # for every draft depth. It is the fallback for configurations that + # have not produced their own progress observation yet. + self.prefix_survival_by_depth: List[Optional[_EMAValue]] = [None] * self.max_draft_step + + # The current verify width is bounded by the proposal built last time. + self.pre_draft_step = self.max_draft_step + + def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> SpecDecodePlan: + pre_draft_step = self.pre_draft_step + if req_num == 0: + self.pre_draft_step = self.max_draft_step + return SpecDecodePlan( + dynamic_batch_size=0, + draft_step=self.max_draft_step, + pre_draft_step=pre_draft_step, + ) + + # A request entering its first decode has only the guaranteed target + # row. Existing requests additionally expose pre_draft_step candidates. + # This bounds Verify by the candidate pool that physically exists. + available_batch_size = req_num + proposal_req_num * pre_draft_step + max_batch_size = min(original_batch_size, available_batch_size) + all_reqs_have_proposals = proposal_req_num == req_num + + if ( + not self.target_infer_costs.has_data() + or not self.draft_infer_costs.has_data() + or not self.progress_ema_by_config + ): + # Costs come from CUDA Graph capture, while progress requires one + # real verification batch. Keep the available proposal intact until + # both parts of J(N, B, d) have an observation. + self.pre_draft_step = self.max_draft_step + return SpecDecodePlan( + dynamic_batch_size=max_batch_size, + draft_step=self.max_draft_step, + pre_draft_step=pre_draft_step, + all_reqs_have_proposals=all_reqs_have_proposals, + ) + + min_batch_size = req_num + batch_sizes = self.target_infer_costs.get_batch_size_keys_between(min_batch_size, max_batch_size) + costs = [ + self._get_cost_ms( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=pre_draft_step, + ) + for dynamic_batch_size in batch_sizes + ] + dynamic_batch_size = batch_sizes[np.argmin(costs)] + draft_step = self._select_draft_step(req_num=req_num) + + self.pre_draft_step = draft_step + return SpecDecodePlan( + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + pre_draft_step=pre_draft_step, + all_reqs_have_proposals=all_reqs_have_proposals, + ) + + def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: + cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs + cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + + def update_feedback( + self, + plan: SpecDecodePlan, + req_num: int, + accept_lengths, + schedule_scores=None, + ) -> None: + # The progress EMA records one complete-batch sample for a single + # (N, B, d) configuration. all_reqs_have_proposals is false if any request is + # on its first decode and therefore has no preceding proposal; its + # structural accept_len=1 would otherwise bias that configuration's + # progress downward. Skip the mixed batch instead of partially + # updating the batch-level statistic with different N/B semantics. + if not plan.all_reqs_have_proposals: + return + self.update_verified_batch( + accept_lengths=accept_lengths, + req_num=req_num, + dynamic_batch_size=plan.dynamic_batch_size, + verified_draft_step=plan.pre_draft_step, + ) + + def update_verified_batch( + self, + accept_lengths, + req_num: int, + dynamic_batch_size: int, + verified_draft_step: int, + ) -> None: + """Record one batch-level progress sample for the verified configuration.""" + + accept_lengths = np.asarray(accept_lengths) + if accept_lengths.size == 0: + return + + config = (int(req_num), int(dynamic_batch_size), int(verified_draft_step)) + progress = float(accept_lengths.sum()) / dynamic_batch_size + if config not in self.progress_ema_by_config: + self.progress_ema_by_config[config] = _EMAValue( + decay=0.9, + init_value=progress, + ) + self.progress_ema_by_config[config].update(progress) + + if dynamic_batch_size == req_num * (verified_draft_step + 1): + for depth in range(1, verified_draft_step + 1): + survival = float(np.mean(accept_lengths > depth)) + survival_ema = self.prefix_survival_by_depth[depth - 1] + if survival_ema is None: + survival_ema = _EMAValue(decay=0.9, init_value=survival) + self.prefix_survival_by_depth[depth - 1] = survival_ema + survival_ema.update(survival) + + def _select_draft_step(self, req_num: int) -> int: + """Choose the best self-consistent next proposal depth. + + The current target batch is bounded by ``pre_draft_step``, but the + proposal generated now is consumed by the next iteration. Evaluate + each candidate depth with the verify widths that depth can create so a + short current proposal cannot permanently prevent deeper drafting. + """ + + best_cost_ms = float("inf") + best_draft_step = self.draft_steps[0] + for draft_step in self.draft_steps: + max_batch_size = req_num * (draft_step + 1) + batch_sizes = self.target_infer_costs.get_batch_size_keys_between(req_num, max_batch_size) + for dynamic_batch_size in batch_sizes: + cost_ms = self._get_cost_ms( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + ) + if cost_ms < best_cost_ms: + best_cost_ms = cost_ms + best_draft_step = draft_step + return best_draft_step + + def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: + """Estimate milliseconds per committed token for one configuration.""" + + accept_ratio = self._estimate_progress( + req_num=req_num, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + ) + total_time = self.target_infer_costs.get(dynamic_batch_size) + self.draft_cost_provider.get_draft_cost_ms( + draft_infer_costs=self.draft_infer_costs, + req_num=req_num, + verify_batch_size=dynamic_batch_size, + draft_step=draft_step, + ) + token_num = min((dynamic_batch_size * accept_ratio), req_num * (draft_step + 1)) + token_num = max(token_num, req_num) + cost_ms = total_time / token_num + return cost_ms + + def _estimate_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: + config = (int(req_num), int(dynamic_batch_size), int(draft_step)) + ema = self.progress_ema_by_config.get(config) + return ema.get() if ema is not None else self._estimate_prefix_progress(*config) + + def _estimate_prefix_progress(self, req_num: int, dynamic_batch_size: int, draft_step: int) -> float: + if not any(self.prefix_survival_by_depth): + return 1.0 + + remaining_draft_rows = dynamic_batch_size - req_num + expected_tokens = float(req_num) + for depth in range(min(draft_step, self.max_draft_step)): + selected_rows = min(req_num, remaining_draft_rows) + if selected_rows <= 0: + break + survival_ema = self.prefix_survival_by_depth[depth] + survival = 1.0 if survival_ema is None else survival_ema.get() + expected_tokens += selected_rows * survival + remaining_draft_rows -= selected_rows + return expected_tokens / dynamic_batch_size diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 943628377e..6cba9bced0 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -9,8 +9,8 @@ FixedSpecPlanner, LightSpecPlanner, SpecDecodePlan, - _InferCostMsTable, ) +from lightllm.server.router.model_infer.mtp_speculative.planner.base import _InferCostMsTable from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer From cec862414266ff7a9e0f374989f2cda8062746dc Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 07:01:41 +0000 Subject: [PATCH 039/103] refactor: pass decode requests to planners --- .../model_infer/mtp_speculative/engine.py | 16 +++------ .../mtp_speculative/planner/dspark.py | 3 +- .../mtp_speculative/planner/fixed.py | 4 ++- .../mtp_speculative/planner/lightspec.py | 8 +++-- .../mtp_speculative/test_planner.py | 35 +++++++++++-------- 5 files changed, 35 insertions(+), 31 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index b51e09b957..488f149148 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -79,18 +79,10 @@ def build_draft_state_from_prefill_overlap( def plan_decode(self, model_input: ModelInput, decode_reqs: List) -> SpecDecodePlan: """Return the fixed or dynamic speculative plan for one decode iteration.""" - req_num = len(decode_reqs) - if isinstance(self.planner, LightSpecPlanner): - # Prefill initializes draft cache but does not create a proposal. - # After one decode, cur_output_len > 1 and the request owns the - # previous iteration's candidates. - proposal_req_num = sum(req.cur_output_len > 1 for req in decode_reqs) - return self.planner.plan( - req_num=req_num, - original_batch_size=model_input.batch_size, - proposal_req_num=proposal_req_num, - ) - return self.planner.plan(req_num=req_num, original_batch_size=model_input.batch_size) + return self.planner.plan( + decode_reqs=decode_reqs, + original_batch_size=model_input.batch_size, + ) def prepare_decode_model_input( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index b68211d709..35e9c4fc74 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -25,7 +25,8 @@ def __init__(self, max_draft_step: int, draft_cost_provider: "BaseSpecProposer") self.draft_infer_costs = _InferCostMsTable() self._pending_verify_batch_sizes = deque(maxlen=2) - def plan(self, req_num: int, original_batch_size: int) -> SpecDecodePlan: + def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: + req_num = len(decode_reqs) if req_num == 0: return SpecDecodePlan( dynamic_batch_size=0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py index 177347638e..a3d76015ea 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py @@ -1,5 +1,7 @@ from __future__ import annotations +from typing import List + from lightllm.server.router.model_infer.mtp_speculative.planner.base import SpecDecodePlan @@ -9,7 +11,7 @@ class FixedSpecPlanner: def __init__(self, max_draft_step: int) -> None: self.max_draft_step = int(max_draft_step) - def plan(self, req_num: int | None = None, original_batch_size: int | None = None) -> SpecDecodePlan: + def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: return SpecDecodePlan( dynamic_batch_size=None, draft_step=self.max_draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index 288d76a0b2..a7449f82e4 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -60,7 +60,8 @@ def __init__( # The current verify width is bounded by the proposal built last time. self.pre_draft_step = self.max_draft_step - def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> SpecDecodePlan: + def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: + req_num = len(decode_reqs) pre_draft_step = self.pre_draft_step if req_num == 0: self.pre_draft_step = self.max_draft_step @@ -73,9 +74,10 @@ def plan(self, req_num: int, original_batch_size: int, proposal_req_num: int) -> # A request entering its first decode has only the guaranteed target # row. Existing requests additionally expose pre_draft_step candidates. # This bounds Verify by the candidate pool that physically exists. - available_batch_size = req_num + proposal_req_num * pre_draft_step + req_num_with_proposals = sum(req.cur_output_len > 1 for req in decode_reqs) + available_batch_size = req_num + req_num_with_proposals * pre_draft_step max_batch_size = min(original_batch_size, available_batch_size) - all_reqs_have_proposals = proposal_req_num == req_num + all_reqs_have_proposals = req_num_with_proposals == req_num if ( not self.target_infer_costs.has_data() diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 6cba9bced0..d5ff573eb0 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -57,6 +57,14 @@ def build_dspark_planner(max_draft_step: int = 3, block_size: int = 3): ) +def build_decode_reqs(req_num: int, req_num_with_proposals: int | None = None): + if req_num_with_proposals is None: + req_num_with_proposals = req_num + return [SimpleNamespace(cur_output_len=2)] * req_num_with_proposals + [ + SimpleNamespace(cur_output_len=1) + ] * (req_num - req_num_with_proposals) + + def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): engine = SpecEngine.__new__(SpecEngine) engine.spec_mode = spec_mode @@ -75,7 +83,7 @@ def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): def test_fixed_planner_returns_static_plan(): - plan = FixedSpecPlanner(max_draft_step=3).plan(req_num=4, original_batch_size=16) + plan = FixedSpecPlanner(max_draft_step=3).plan(decode_reqs=build_decode_reqs(4), original_batch_size=16) assert not plan.is_dynamic assert plan.dynamic_batch_size is None @@ -187,9 +195,8 @@ def update_mtp_verify_step_num(self, verify_step_num: int): def test_lightspec_stays_full_width_until_costs_are_profiled(): plan = build_lightspec_planner().plan( - req_num=2, + decode_reqs=build_decode_reqs(2), original_batch_size=8, - proposal_req_num=2, ) assert plan.dynamic_batch_size == 8 @@ -202,7 +209,7 @@ def test_lightspec_collects_full_width_progress_before_adapting(): planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=False) planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) - plan = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) assert plan.dynamic_batch_size == 8 assert plan.draft_step == plan.pre_draft_step == 3 @@ -222,7 +229,7 @@ def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): verified_draft_step=3, ) - plan = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) assert plan.dynamic_batch_size == 4 assert plan.pre_draft_step == 3 @@ -246,7 +253,7 @@ def test_lightspec_compacts_block_verify_without_changing_draft_shape(): verified_draft_step=7, ) - plan = planner.plan(req_num=2, original_batch_size=16, proposal_req_num=2) + plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=16) assert plan.dynamic_batch_size < 16 assert plan.draft_step == plan.pre_draft_step == 7 @@ -255,9 +262,9 @@ def test_lightspec_compacts_block_verify_without_changing_draft_shape(): def test_lightspec_bounds_verify_to_existing_proposals(): planner = build_lightspec_planner() - cold_start = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=0) - mixed_batch = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=1) - ready_batch = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + cold_start = planner.plan(decode_reqs=build_decode_reqs(2, 0), original_batch_size=8) + mixed_batch = planner.plan(decode_reqs=build_decode_reqs(2, 1), original_batch_size=8) + ready_batch = planner.plan(decode_reqs=build_decode_reqs(2, 2), original_batch_size=8) assert cold_start.dynamic_batch_size == 2 assert mixed_batch.dynamic_batch_size == 5 @@ -267,7 +274,7 @@ def test_lightspec_bounds_verify_to_existing_proposals(): assert ready_batch.all_reqs_have_proposals -def test_engine_counts_requests_with_a_previous_proposal(): +def test_engine_lets_planner_count_requests_with_a_previous_proposal(): engine = SpecEngine.__new__(SpecEngine) engine.planner = build_lightspec_planner() model_input = SimpleNamespace(batch_size=8) @@ -299,7 +306,7 @@ def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): verified_draft_step=3, ) - plan = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) assert planner.draft_steps == (1, 2, 3) assert plan.draft_step >= 1 @@ -444,7 +451,7 @@ def test_lightspec_short_current_proposal_can_recover_to_a_deeper_draft(): ) planner.pre_draft_step = 1 - plan = planner.plan(req_num=2, original_batch_size=8, proposal_req_num=2) + plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) assert plan.dynamic_batch_size <= 4 assert plan.pre_draft_step == 1 @@ -497,14 +504,14 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): req_num=2, accept_lengths_cpu=torch.tensor([1, 1], dtype=torch.int32), ) - first_plan = planner.plan(req_num=2, original_batch_size=8) + first_plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) engine.update_planner_feedback( plan=plan, proposal=proposal, req_num=2, accept_lengths_cpu=torch.tensor([1, 1], dtype=torch.int32), ) - second_plan = planner.plan(req_num=2, original_batch_size=8) + second_plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) assert first_plan.dynamic_batch_size == 8 assert second_plan.dynamic_batch_size == 4 From 5963199f4013cfdf6b3ab7759e1fe13a4ecb927b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 07:10:45 +0000 Subject: [PATCH 040/103] refactor: move draft step selection to planner --- .../router/model_infer/mtp_speculative/engine.py | 1 + .../mtp_speculative/planner/lightspec.py | 15 ++++++++++++++- .../model_infer/mtp_speculative/proposers/base.py | 7 +------ .../mtp_speculative/proposers/eagle_mtp.py | 3 --- .../mtp_speculative/proposers/parallel_block.py | 3 --- .../mtp_speculative/proposers/vanilla_mtp.py | 3 --- .../model_infer/mtp_speculative/test_planner.py | 11 +++++++++++ 7 files changed, 27 insertions(+), 16 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 488f149148..0ec63358e0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -307,6 +307,7 @@ def _build_decode_planner(self): ) return LightSpecPlanner( + spec_mode=self.spec_mode, max_draft_step=self.backend.max_draft_step, draft_cost_provider=self.proposer, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index a7449f82e4..74fbfad41e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -38,12 +38,14 @@ class LightSpecPlanner: def __init__( self, + spec_mode: str, max_draft_step: int, draft_cost_provider: "BaseSpecProposer", ) -> None: + self.spec_mode = spec_mode self.max_draft_step = int(max_draft_step) self.draft_cost_provider = draft_cost_provider - self.draft_steps = draft_cost_provider.get_draft_steps() + self.draft_steps = self.get_draft_steps() self.target_infer_costs = _InferCostMsTable() self.draft_infer_costs = _InferCostMsTable() @@ -60,6 +62,17 @@ def __init__( # The current verify width is bounded by the proposal built last time. self.pre_draft_step = self.max_draft_step + def get_draft_steps(self) -> Tuple[int, ...]: + """Return the draft configurations supported by the current MTP mode.""" + + if self.spec_mode in ("vanilla_with_att", "vanilla_no_att"): + return tuple(range(self.max_draft_step + 1)) + if self.spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3"): + return tuple(range(1, self.max_draft_step + 1)) + if self.spec_mode == "dflash": + return (self.max_draft_step,) + raise ValueError(f"unsupported LightSpec mode: {self.spec_mode}") + def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: req_num = len(decode_reqs) pre_draft_step = self.pre_draft_step diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 6f9d9e9145..20735b109b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -1,7 +1,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Optional, Tuple +from typing import TYPE_CHECKING, Optional import torch @@ -45,11 +45,6 @@ def __init__(self, *, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> None: self.backend = backend self.enable_dynmaic_mtp = bool(enable_dynmaic_mtp) - def get_draft_steps(self) -> Tuple[int, ...]: - """Return the draft configurations supported by this proposer.""" - - raise NotImplementedError - def get_draft_cost_ms( self, draft_infer_costs, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py index cb5c9ce299..fd63455f9b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py @@ -11,9 +11,6 @@ class AutoregressiveEagleProposer(BaseSpecProposer): """Shared autoregressive drafting flow for EAGLE-family proposers.""" - def get_draft_steps(self): - return tuple(range(1, self.backend.max_draft_step + 1)) - def get_draft_cost_ms( self, draft_infer_costs, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py index 80a7ce41be..07f9694693 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py @@ -16,9 +16,6 @@ class ParallelBlockProposer(BaseSpecProposer): backbone forward. """ - def get_draft_steps(self): - return (self.backend.max_draft_step,) - def get_draft_cost_ms( self, draft_infer_costs, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py index a5fad08c84..8c17ed8501 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py @@ -13,9 +13,6 @@ class VanillaMTPProposer(BaseSpecProposer): hidden state produced by module i - 1 and predicts the next candidate. """ - def get_draft_steps(self): - return tuple(range(self.backend.max_draft_step + 1)) - def get_draft_cost_ms( self, draft_infer_costs, diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index d5ff573eb0..8b18e9e612 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -36,7 +36,14 @@ def build_lightspec_planner( proposer_class=VanillaMTPProposer, block_size: int = 3, ): + spec_mode = { + VanillaMTPProposer: "vanilla_with_att", + EagleMTPProposer: "eagle_with_att", + Eagle3Proposer: "eagle3", + DFlashProposer: "dflash", + }[proposer_class] return LightSpecPlanner( + spec_mode=spec_mode, max_draft_step=max_draft_step, draft_cost_provider=build_draft_cost_provider( proposer_class=proposer_class, @@ -103,6 +110,10 @@ def test_engine_routes_only_dspark_to_the_confidence_planner(): assert isinstance(build_planner("eagle3", enable_dynmaic_mtp=False), FixedSpecPlanner) assert isinstance(build_planner("dspark"), DSparkPlanner) + vanilla_planner = build_planner("vanilla_with_att") + assert isinstance(vanilla_planner, LightSpecPlanner) + assert vanilla_planner.draft_steps == (0, 1, 2, 3) + dflash_planner = build_planner("dflash") assert isinstance(dflash_planner, LightSpecPlanner) assert dflash_planner.draft_steps == (3,) From 6e25ddc2ac0dca42ad09e90a730431d8973db241 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Tue, 18 Aug 2026 09:47:58 +0000 Subject: [PATCH 041/103] refactor: move MTP cost modeling to planners --- .../model_infer/mtp_speculative/engine.py | 33 +-- .../mtp_speculative/planner/dspark.py | 38 ++-- .../mtp_speculative/planner/lightspec.py | 62 ++++-- .../mtp_speculative/proposers/base.py | 11 - .../mtp_speculative/proposers/eagle_mtp.py | 20 +- .../proposers/parallel_block.py | 14 -- .../mtp_speculative/proposers/vanilla_mtp.py | 10 - .../mtp_speculative/test_planner.py | 204 +++++++++++------- 8 files changed, 204 insertions(+), 188 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 0ec63358e0..1301b2470e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -40,7 +40,6 @@ def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: enable_dynmaic_mtp=enable_dynmaic_mtp, ) self.planner = self._build_decode_planner() - self._register_cuda_graph_costs() # Prefill draft-state initialization. @@ -301,37 +300,9 @@ def _build_decode_planner(self): if not self.enable_dynmaic_mtp: return FixedSpecPlanner(max_draft_step=self.backend.max_draft_step) if self.spec_mode == "dspark": - return DSparkPlanner( - max_draft_step=self.backend.max_draft_step, - draft_cost_provider=self.proposer, - ) + return DSparkPlanner(backend=self.backend) return LightSpecPlanner( spec_mode=self.spec_mode, - max_draft_step=self.backend.max_draft_step, - draft_cost_provider=self.proposer, + backend=self.backend, ) - - def _register_cuda_graph_costs(self) -> None: - if not self.enable_dynmaic_mtp: - return - - target_graph = self.backend.model.graph - if target_graph is not None: - for batch_size, infer_cost_ms in target_graph.infer_cost_ms_by_batch_size.items(): - self.planner.update_infer_cost( - batch_size=batch_size, - infer_cost_ms=infer_cost_ms, - is_draft_model=False, - ) - - for draft_model in self.backend.draft_models: - draft_graph = draft_model.graph - if draft_graph is None: - continue - for batch_size, infer_cost_ms in draft_graph.infer_cost_ms_by_batch_size.items(): - self.planner.update_infer_cost( - batch_size=batch_size, - infer_cost_ms=infer_cost_ms, - is_draft_model=True, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index 35e9c4fc74..3cc23000fb 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -1,15 +1,12 @@ from __future__ import annotations from collections import deque -from typing import TYPE_CHECKING, Dict, List, Optional +from typing import Dict, List, Optional import numpy as np from lightllm.server.router.model_infer.mtp_speculative.planner.base import SpecDecodePlan, _InferCostMsTable -if TYPE_CHECKING: - from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer - class DSparkPlanner: """DSpark's confidence-based verify-capacity planner. @@ -18,11 +15,13 @@ class DSparkPlanner: the target verify capacity two iterations later. """ - def __init__(self, max_draft_step: int, draft_cost_provider: "BaseSpecProposer") -> None: - self.max_draft_step = int(max_draft_step) - self.draft_cost_provider = draft_cost_provider + def __init__(self, backend) -> None: + self.backend = backend + self.max_draft_step = int(backend.max_draft_step) + self.block_size = int(backend.draft_models[0].block_size) self.target_infer_costs = _InferCostMsTable() self.draft_infer_costs = _InferCostMsTable() + self._register_cuda_graph_costs() self._pending_verify_batch_sizes = deque(maxlen=2) def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: @@ -50,9 +49,25 @@ def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: pre_draft_step=self.max_draft_step, ) - def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: - cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs - cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + def get_draft_cost_ms(self, req_num: int, verify_batch_size: int, draft_step: int) -> float: + """Return the cost of committing verify rows and generating one complete block.""" + + extend_cost_ms = self.draft_infer_costs.estimate(verify_batch_size) + block_cost_ms = self.draft_infer_costs.estimate(req_num * self.block_size) + return extend_cost_ms + block_cost_ms + + def _register_cuda_graph_costs(self) -> None: + target_graph = self.backend.model.graph + if target_graph is not None: + for batch_size, infer_cost_ms in target_graph.infer_cost_ms_by_batch_size.items(): + self.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + + for draft_model in self.backend.draft_models: + draft_graph = draft_model.graph + if draft_graph is None: + continue + for batch_size, infer_cost_ms in draft_graph.infer_cost_ms_by_batch_size.items(): + self.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) def update_feedback( self, @@ -140,8 +155,7 @@ def _select_dynamic_batch_size_from_survival_scores( for dynamic_batch_size in sorted(candidate_batch_sizes): selected_draft_count = dynamic_batch_size - int(req_num) expected_tokens = float(req_num) + float(expected_accepts_by_count[selected_draft_count]) - round_ms = self.target_infer_costs.get(dynamic_batch_size) + self.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=self.draft_infer_costs, + round_ms = self.target_infer_costs.estimate(dynamic_batch_size) + self.get_draft_cost_ms( req_num=req_num, verify_batch_size=dynamic_batch_size, draft_step=self.max_draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index 74fbfad41e..83a719acf5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Dict, List, Optional, Tuple +from typing import Dict, List, Optional, Tuple import numpy as np @@ -10,9 +10,6 @@ _InferCostMsTable, ) -if TYPE_CHECKING: - from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer - class LightSpecPlanner: """Choose the current verify budget and the next draft configuration. @@ -31,24 +28,25 @@ class LightSpecPlanner: / B`` per iteration, then smoothed with an EMA. It is never updated once per request; doing so would make adaptation depend on concurrency. - The proposer supplies its valid draft configurations and complete draft - cost. The planner therefore stays independent of the proposal algorithm and - its physical execution layout. + The current MTP mode determines both its valid draft configurations and + complete draft cost. The planner owns those scheduling inputs instead of + depending on the proposer implementation. """ def __init__( self, spec_mode: str, - max_draft_step: int, - draft_cost_provider: "BaseSpecProposer", + backend, ) -> None: self.spec_mode = spec_mode - self.max_draft_step = int(max_draft_step) - self.draft_cost_provider = draft_cost_provider + self.backend = backend + self.max_draft_step = int(backend.max_draft_step) + self.block_size = int(backend.draft_models[0].block_size) if self.spec_mode == "dflash" else None self.draft_steps = self.get_draft_steps() self.target_infer_costs = _InferCostMsTable() self.draft_infer_costs = _InferCostMsTable() + self._register_cuda_graph_costs() # Each observation is the normalized committed progress U / B from one # complete batch. The draft configuration is part of the key because a @@ -65,14 +63,45 @@ def __init__( def get_draft_steps(self) -> Tuple[int, ...]: """Return the draft configurations supported by the current MTP mode.""" - if self.spec_mode in ("vanilla_with_att", "vanilla_no_att"): + if self.spec_mode in ("vanilla_no_att", "eagle_no_att"): return tuple(range(self.max_draft_step + 1)) - if self.spec_mode in ("eagle_with_att", "eagle_no_att", "eagle3"): + if self.spec_mode in ("vanilla_with_att", "eagle_with_att", "eagle3"): return tuple(range(1, self.max_draft_step + 1)) if self.spec_mode == "dflash": return (self.max_draft_step,) raise ValueError(f"unsupported LightSpec mode: {self.spec_mode}") + def get_draft_cost_ms(self, req_num: int, verify_batch_size: int, draft_step: int) -> float: + """Return the complete draft cost for one ``(N, B, d)`` configuration.""" + + if self.spec_mode in ("vanilla_no_att", "eagle_no_att"): + return self.draft_infer_costs.estimate(req_num) * draft_step + if self.spec_mode in ("vanilla_with_att", "eagle_with_att", "eagle3"): + assert draft_step > 0, f"{self.spec_mode} requires draft_step to be greater than 0" + draft_cost_ms = self.draft_infer_costs.estimate(verify_batch_size) + if draft_step > 1: + draft_cost_ms += self.draft_infer_costs.estimate(req_num) * (draft_step - 1) + return draft_cost_ms + if self.spec_mode == "dflash": + assert self.block_size is not None + extend_cost_ms = self.draft_infer_costs.estimate(verify_batch_size) + block_cost_ms = self.draft_infer_costs.estimate(req_num * self.block_size) + return extend_cost_ms + block_cost_ms + raise ValueError(f"unsupported LightSpec mode: {self.spec_mode}") + + def _register_cuda_graph_costs(self) -> None: + target_graph = self.backend.model.graph + if target_graph is not None: + for batch_size, infer_cost_ms in target_graph.infer_cost_ms_by_batch_size.items(): + self.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + + for draft_model in self.backend.draft_models: + draft_graph = draft_model.graph + if draft_graph is None: + continue + for batch_size, infer_cost_ms in draft_graph.infer_cost_ms_by_batch_size.items(): + self.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: req_num = len(decode_reqs) pre_draft_step = self.pre_draft_step @@ -129,10 +158,6 @@ def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: all_reqs_have_proposals=all_reqs_have_proposals, ) - def update_infer_cost(self, batch_size: int, infer_cost_ms: float, is_draft_model: bool) -> None: - cost_table = self.draft_infer_costs if is_draft_model else self.target_infer_costs - cost_table.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - def update_feedback( self, plan: SpecDecodePlan, @@ -219,8 +244,7 @@ def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) - dynamic_batch_size=dynamic_batch_size, draft_step=draft_step, ) - total_time = self.target_infer_costs.get(dynamic_batch_size) + self.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=self.draft_infer_costs, + total_time = self.target_infer_costs.estimate(dynamic_batch_size) + self.get_draft_cost_ms( req_num=req_num, verify_batch_size=dynamic_batch_size, draft_step=draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 20735b109b..1d275ea375 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -45,17 +45,6 @@ def __init__(self, *, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> None: self.backend = backend self.enable_dynmaic_mtp = bool(enable_dynmaic_mtp) - def get_draft_cost_ms( - self, - draft_infer_costs, - req_num: int, - verify_batch_size: int, - draft_step: int, - ) -> float: - """Return the complete draft cost for one ``(N, B, d)`` configuration.""" - - raise NotImplementedError - def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: """Allocate draft-owned temporary KV slots.""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py index fd63455f9b..96c1b74699 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py @@ -11,18 +11,6 @@ class AutoregressiveEagleProposer(BaseSpecProposer): """Shared autoregressive drafting flow for EAGLE-family proposers.""" - def get_draft_cost_ms( - self, - draft_infer_costs, - req_num: int, - verify_batch_size: int, - draft_step: int, - ) -> float: - draft_cost_ms = draft_infer_costs.estimate(verify_batch_size) - if draft_step > 1: - draft_cost_ms += draft_infer_costs.get(req_num) * (draft_step - 1) - return draft_cost_ms - def build_draft_state_from_prefill( self, target_model_input: ModelInput, @@ -118,7 +106,6 @@ def propose_next( verify_row_count = int(next_token_ids.shape[0]) request_count = int(b_req_mtp_start_loc.shape[0]) - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() proposal_token_ids = next_token_ids.new_full( (verify_row_count, draft_step + 1), fill_value=1, @@ -134,7 +121,14 @@ def propose_next( if collect_schedule_scores else None ) + if draft_step == 0: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + schedule_scores=schedule_scores, + ) + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() draft_model = self.backend.draft_models[0] position_delta = main_model_input.b_position_delta self.prepare_verify_extend_input( diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py index 07f9694693..859a8f18a2 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py @@ -16,20 +16,6 @@ class ParallelBlockProposer(BaseSpecProposer): backbone forward. """ - def get_draft_cost_ms( - self, - draft_infer_costs, - req_num: int, - verify_batch_size: int, - draft_step: int, - ) -> float: - block_size = self.backend.draft_models[0].block_size - # One forward commits verified rows; another generates one complete - # checkpoint-defined block per request. - extend_cost_ms = draft_infer_costs.estimate(verify_batch_size) - block_cost_ms = draft_infer_costs.get(req_num * block_size) - return extend_cost_ms + block_cost_ms - @torch.no_grad() def build_draft_state_from_prefill( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py index 8c17ed8501..a5a6f309b1 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py @@ -13,16 +13,6 @@ class VanillaMTPProposer(BaseSpecProposer): hidden state produced by module i - 1 and predicts the next candidate. """ - def get_draft_cost_ms( - self, - draft_infer_costs, - req_num: int, - verify_batch_size: int, - draft_step: int, - ) -> float: - # Every selected MTP module processes the complete verify batch. - return draft_infer_costs.get(verify_batch_size) * draft_step - def build_draft_state_from_prefill( self, target_model_input: ModelInput, diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 8b18e9e612..38eb3e0386 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -1,6 +1,7 @@ from types import SimpleNamespace import numpy as np +import pytest import torch from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine @@ -23,45 +24,29 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_mtp import VanillaMTPProposer -def build_draft_cost_provider(proposer_class, max_draft_step: int = 3, block_size: int = 3): - backend = SimpleNamespace( - max_draft_step=max_draft_step, - draft_models=[SimpleNamespace(block_size=block_size)], - ) - return proposer_class(backend=backend, enable_dynmaic_mtp=True) - - def build_lightspec_planner( max_draft_step: int = 3, - proposer_class=VanillaMTPProposer, + spec_mode: str = "vanilla_with_att", block_size: int = 3, ): - spec_mode = { - VanillaMTPProposer: "vanilla_with_att", - EagleMTPProposer: "eagle_with_att", - Eagle3Proposer: "eagle3", - DFlashProposer: "dflash", - }[proposer_class] + backend = SimpleNamespace( + max_draft_step=max_draft_step, + model=SimpleNamespace(graph=None), + draft_models=[SimpleNamespace(block_size=block_size, graph=None)], + ) return LightSpecPlanner( spec_mode=spec_mode, - max_draft_step=max_draft_step, - draft_cost_provider=build_draft_cost_provider( - proposer_class=proposer_class, - max_draft_step=max_draft_step, - block_size=block_size, - ), + backend=backend, ) def build_dspark_planner(max_draft_step: int = 3, block_size: int = 3): - return DSparkPlanner( + backend = SimpleNamespace( max_draft_step=max_draft_step, - draft_cost_provider=build_draft_cost_provider( - proposer_class=DSparkProposer, - max_draft_step=max_draft_step, - block_size=block_size, - ), + model=SimpleNamespace(graph=None), + draft_models=[SimpleNamespace(block_size=block_size, graph=None)], ) + return DSparkPlanner(backend=backend) def build_decode_reqs(req_num: int, req_num_with_proposals: int | None = None): @@ -78,14 +63,9 @@ def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): engine.enable_dynmaic_mtp = enable_dynmaic_mtp engine.backend = SimpleNamespace( max_draft_step=3, - draft_models=[SimpleNamespace(block_size=3)], - ) - proposer_class = { - "dspark": DSparkProposer, - "dflash": DFlashProposer, - "eagle3": Eagle3Proposer, - }.get(spec_mode, VanillaMTPProposer) - engine.proposer = build_draft_cost_provider(proposer_class) + model=SimpleNamespace(graph=None), + draft_models=[SimpleNamespace(block_size=3, graph=None)], + ) return engine._build_decode_planner() @@ -112,7 +92,13 @@ def test_engine_routes_only_dspark_to_the_confidence_planner(): vanilla_planner = build_planner("vanilla_with_att") assert isinstance(vanilla_planner, LightSpecPlanner) - assert vanilla_planner.draft_steps == (0, 1, 2, 3) + assert vanilla_planner.draft_steps == (1, 2, 3) + + vanilla_no_att_planner = build_planner("vanilla_no_att") + assert vanilla_no_att_planner.draft_steps == (0, 1, 2, 3) + + eagle_no_att_planner = build_planner("eagle_no_att") + assert eagle_no_att_planner.draft_steps == (0, 1, 2, 3) dflash_planner = build_planner("dflash") assert isinstance(dflash_planner, LightSpecPlanner) @@ -123,6 +109,24 @@ def test_engine_routes_only_dspark_to_the_confidence_planner(): assert eagle_planner.draft_steps == (1, 2, 3) +def test_dynamic_planner_registers_cuda_graph_costs_from_backend(): + backend = SimpleNamespace( + max_draft_step=3, + model=SimpleNamespace(graph=SimpleNamespace(infer_cost_ms_by_batch_size={4: 1.2})), + draft_models=[ + SimpleNamespace( + block_size=3, + graph=SimpleNamespace(infer_cost_ms_by_batch_size={4: 0.3}), + ) + ], + ) + + planner = LightSpecPlanner(spec_mode="vanilla_with_att", backend=backend) + + assert planner.target_infer_costs.estimate(4) == 1.2 + assert planner.draft_infer_costs.estimate(4) == 0.3 + + def test_proposer_families_use_drafter_standard_abstractions(): assert issubclass(EagleMTPProposer, AutoregressiveEagleProposer) assert issubclass(Eagle3Proposer, AutoregressiveEagleProposer) @@ -131,6 +135,25 @@ def test_proposer_families_use_drafter_standard_abstractions(): assert not issubclass(DSparkProposer, DFlashProposer) +def test_eagle_proposer_skips_draft_forward_for_zero_steps(): + proposer = EagleMTPProposer( + backend=SimpleNamespace(draft_models=[]), + enable_dynmaic_mtp=True, + ) + next_token_ids = torch.tensor([10, 11], dtype=torch.int64) + + proposal = proposer.propose_next( + main_model_input=None, + main_model_output=None, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=torch.tensor([0, 1], dtype=torch.int32), + draft_step=0, + ) + + assert proposal.token_ids.tolist() == [[10], [11]] + assert proposal.schedule_scores.shape == (2, 0) + + def test_dynamic_plan_filters_selected_rows(): plan = SpecDecodePlan(dynamic_batch_size=2, draft_step=3, pre_draft_step=3) reqs = ["req0", "req0", "req1", "req1"] @@ -215,10 +238,10 @@ def test_lightspec_stays_full_width_until_costs_are_profiled(): def test_lightspec_collects_full_width_progress_before_adapting(): - planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + planner = build_lightspec_planner(spec_mode="eagle_with_att") for batch_size in (2, 4, 8): - planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=False) - planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) + planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) @@ -227,11 +250,11 @@ def test_lightspec_collects_full_width_progress_before_adapting(): def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): - planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + planner = build_lightspec_planner(spec_mode="eagle_with_att") for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 10.0)): - planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) for batch_size, draft_cost in ((2, 0.1), (4, 0.1), (8, 0.2)): - planner.update_infer_cost(batch_size, draft_cost, is_draft_model=True) + planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=draft_cost) for _ in range(8): planner.update_verified_batch( accept_lengths=[4, 4], @@ -250,12 +273,12 @@ def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): def test_lightspec_compacts_block_verify_without_changing_draft_shape(): planner = build_lightspec_planner( max_draft_step=7, - proposer_class=DFlashProposer, + spec_mode="dflash", block_size=7, ) for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 3.0), (16, 8.0)): - planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) - planner.update_infer_cost(batch_size=14, infer_cost_ms=0.5, is_draft_model=True) + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) + planner.draft_infer_costs.update(batch_size=14, infer_cost_ms=0.5) for _ in range(8): planner.update_verified_batch( accept_lengths=[3, 3], @@ -306,10 +329,10 @@ def test_engine_lets_planner_count_requests_with_a_previous_proposal(): def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): - planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + planner = build_lightspec_planner(spec_mode="eagle_with_att") for batch_size in (2, 4, 8): - planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=False) - planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) + planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) planner.update_verified_batch( accept_lengths=[2, 2], req_num=2, @@ -323,26 +346,40 @@ def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): assert plan.draft_step >= 1 -def test_vanilla_cost_provider_prices_each_selected_mtp_module(): - planner = build_lightspec_planner() - planner.update_infer_cost(batch_size=8, infer_cost_ms=0.25, is_draft_model=True) +def test_vanilla_with_attention_planner_prices_extend_then_normal_batches(): + planner = build_lightspec_planner(spec_mode="vanilla_with_att") + planner.draft_infer_costs.update(batch_size=2, infer_cost_ms=0.1) - draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=planner.draft_infer_costs, - req_num=2, + draft_cost_ms = planner.get_draft_cost_ms( + req_num=4, verify_batch_size=8, draft_step=3, ) - assert draft_cost_ms == 0.75 + assert np.isclose(draft_cost_ms, 0.8) + with pytest.raises(AssertionError, match="requires draft_step to be greater than 0"): + planner.get_draft_cost_ms(req_num=4, verify_batch_size=8, draft_step=0) + + +def test_no_attention_planner_prices_only_normal_request_batches(): + for spec_mode in ("vanilla_no_att", "eagle_no_att"): + planner = build_lightspec_planner(spec_mode=spec_mode) + planner.draft_infer_costs.update(batch_size=2, infer_cost_ms=0.1) + + draft_cost_ms = planner.get_draft_cost_ms( + req_num=4, + verify_batch_size=8, + draft_step=3, + ) + + assert np.isclose(draft_cost_ms, 0.6) -def test_eagle_cost_provider_prices_extend_and_decode_separately(): - planner = build_lightspec_planner(proposer_class=EagleMTPProposer) - planner.update_infer_cost(batch_size=2, infer_cost_ms=0.25, is_draft_model=True) +def test_eagle_planner_prices_extend_and_decode_separately(): + planner = build_lightspec_planner(spec_mode="eagle_with_att") + planner.draft_infer_costs.update(batch_size=2, infer_cost_ms=0.25) - draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=planner.draft_infer_costs, + draft_cost_ms = planner.get_draft_cost_ms( req_num=2, verify_batch_size=8, draft_step=3, @@ -351,17 +388,16 @@ def test_eagle_cost_provider_prices_extend_and_decode_separately(): assert draft_cost_ms == 1.5 -def test_autoregressive_eagle_cost_provider_prices_extend_and_decode_rows(): - for proposer_class in (EagleMTPProposer, Eagle3Proposer): +def test_autoregressive_eagle_planner_prices_extend_and_decode_rows(): + for spec_mode in ("eagle_with_att", "eagle3"): planner = build_lightspec_planner( max_draft_step=7, - proposer_class=proposer_class, + spec_mode=spec_mode, ) for batch_size in (2, 4, 8, 16): - planner.update_infer_cost(batch_size, infer_cost_ms=float(batch_size), is_draft_model=True) + planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) - draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=planner.draft_infer_costs, + draft_cost_ms = planner.get_draft_cost_ms( req_num=8, verify_batch_size=16, draft_step=7, @@ -370,16 +406,15 @@ def test_autoregressive_eagle_cost_provider_prices_extend_and_decode_rows(): assert draft_cost_ms == 64.0 -def test_block_cost_provider_prices_commit_and_complete_block(): +def test_block_planner_prices_commit_and_complete_block(): planner = build_lightspec_planner( max_draft_step=7, - proposer_class=DFlashProposer, + spec_mode="dflash", block_size=7, ) - planner.update_infer_cost(batch_size=14, infer_cost_ms=0.7, is_draft_model=True) + planner.draft_infer_costs.update(batch_size=8, infer_cost_ms=0.4) - draft_cost_ms = planner.draft_cost_provider.get_draft_cost_ms( - draft_infer_costs=planner.draft_infer_costs, + draft_cost_ms = planner.get_draft_cost_ms( req_num=2, verify_batch_size=16, draft_step=7, @@ -388,6 +423,19 @@ def test_block_cost_provider_prices_commit_and_complete_block(): assert np.isclose(draft_cost_ms, 1.5) +def test_dspark_planner_prices_commit_and_complete_block(): + planner = build_dspark_planner(max_draft_step=3, block_size=3) + planner.draft_infer_costs.update(batch_size=8, infer_cost_ms=0.8) + + draft_cost_ms = planner.get_draft_cost_ms( + req_num=4, + verify_batch_size=8, + draft_step=3, + ) + + assert np.isclose(draft_cost_ms, 2.0) + + def test_lightspec_records_one_batch_observation_per_configuration(): planner = build_lightspec_planner() @@ -450,10 +498,10 @@ def test_lightspec_does_not_transfer_deep_progress_to_short_drafts(): def test_lightspec_short_current_proposal_can_recover_to_a_deeper_draft(): - planner = build_lightspec_planner(proposer_class=EagleMTPProposer) + planner = build_lightspec_planner(spec_mode="eagle_with_att") for batch_size, target_cost in ((2, 1.0), (4, 1.1), (6, 1.2), (8, 1.3)): - planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) - planner.update_infer_cost(batch_size, 0.01, is_draft_model=True) + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) + planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=0.01) planner.update_verified_batch( accept_lengths=[4, 4], req_num=2, @@ -496,8 +544,8 @@ def test_engine_skips_feedback_for_a_mixed_proposal_batch(): def test_dspark_applies_confidence_capacity_after_two_step_delay(): planner = build_dspark_planner() for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 10.0)): - planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) - planner.update_infer_cost(batch_size=6, infer_cost_ms=0.5, is_draft_model=True) + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) + planner.draft_infer_costs.update(batch_size=6, infer_cost_ms=0.5) confidence_probs = np.asarray([[0.9, 0.9, 0.9]] * 2, dtype=np.float64) plan = SpecDecodePlan(dynamic_batch_size=8, draft_step=3, pre_draft_step=3) proposal = SpecProposal( @@ -532,8 +580,8 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): def test_dspark_uses_confidence_scores_without_acceptance_ema(): planner = build_dspark_planner() for batch_size, target_cost in ((2, 1.0), (4, 1.2), (8, 10.0)): - planner.update_infer_cost(batch_size, target_cost, is_draft_model=False) - planner.update_infer_cost(batch_size=6, infer_cost_ms=0.5, is_draft_model=True) + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) + planner.draft_infer_costs.update(batch_size=6, infer_cost_ms=0.5) low_confidence = np.cumprod(np.full((2, 3), 0.1), axis=1) high_confidence = np.cumprod(np.full((2, 3), 0.9), axis=1) From 9a3cb61d7899171f166f2983821363a76cf11147 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 00:27:44 +0000 Subject: [PATCH 042/103] refactor: require CUDA graphs for MTP planning --- lightllm/server/api_start.py | 3 +++ .../model_infer/mtp_speculative/planner/base.py | 3 --- .../mtp_speculative/planner/dspark.py | 15 ++++++--------- .../mtp_speculative/planner/lightspec.py | 16 ++++++++++------ unit_tests/server/test_mtp_start_args.py | 12 ++++++++++++ 5 files changed, 31 insertions(+), 18 deletions(-) create mode 100644 unit_tests/server/test_mtp_start_args.py diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index dc6e999b86..1e949f9eb5 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -36,6 +36,9 @@ def _set_envs_and_config(args: StartArgs): def _launch_subprocesses(args: StartArgs): _set_envs_and_config(args) + if args.mtp_mode is not None: + assert not args.disable_cudagraph, "--disable_cudagraph is not supported when --mtp_mode is enabled" + auto_set_max_req_total_len(args) auto_set_fused_shared_experts(args) set_unique_server_name(args) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py index f6e1304de8..1e8f9cd91b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -48,9 +48,6 @@ def __init__(self) -> None: def update(self, batch_size: int, infer_cost_ms: float) -> None: self.infer_cost_ms_table[int(batch_size)] = float(infer_cost_ms) - def has_data(self) -> bool: - return len(self.infer_cost_ms_table) > 0 - def get(self, batch_size: int) -> float: batch_size = int(batch_size) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index 3cc23000fb..9f55a14300 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -35,13 +35,12 @@ def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: full_batch_size = original_batch_size dynamic_batch_size = full_batch_size - if self.target_infer_costs.has_data() and self.draft_infer_costs.has_data(): - delayed_batch_size = self._pop_delayed_batch_size( - req_num=req_num, - max_batch_size=full_batch_size, - ) - if delayed_batch_size is not None: - dynamic_batch_size = delayed_batch_size + delayed_batch_size = self._pop_delayed_batch_size( + req_num=req_num, + max_batch_size=full_batch_size, + ) + if delayed_batch_size is not None: + dynamic_batch_size = delayed_batch_size return SpecDecodePlan( dynamic_batch_size=dynamic_batch_size, @@ -93,8 +92,6 @@ def update_confidence_probs(self, confidence_probs, req_num: int) -> None: if req_num <= 0: return - if not self.target_infer_costs.has_data() or not self.draft_infer_costs.has_data(): - return probs = np.asarray(confidence_probs, dtype=np.float64) if probs.ndim != 2 or probs.shape[1] == 0: diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index 83a719acf5..1b82285a3d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -121,14 +121,10 @@ def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: max_batch_size = min(original_batch_size, available_batch_size) all_reqs_have_proposals = req_num_with_proposals == req_num - if ( - not self.target_infer_costs.has_data() - or not self.draft_infer_costs.has_data() - or not self.progress_ema_by_config - ): + if not self.progress_ema_by_config: # Costs come from CUDA Graph capture, while progress requires one # real verification batch. Keep the available proposal intact until - # both parts of J(N, B, d) have an observation. + # J(N, B, d) has a progress observation. self.pre_draft_step = self.max_draft_step return SpecDecodePlan( dynamic_batch_size=max_batch_size, @@ -202,6 +198,14 @@ def update_verified_batch( ) self.progress_ema_by_config[config].update(progress) + # 只有完整 verify 布局 B = N * (d + 1) 才能统计各 draft 深度的前缀存活率。 + # 如果 B 被动态压缩,较深位置的候选可能根本没有参与验证;此时把这些位置 + # 当作未接受会系统性低估深层候选的效果。 + # + # accept_lengths 包含每个请求必然提交的 1 个 target token,因此 + # accept_lengths > depth 表示该请求至少接受了前 depth 个 draft token。 + # survival 是第 depth 个 draft token 的前缀存活概率,并通过 EMA 平滑; + # planner 使用它估算尚未实际运行过的 (N, B, d) 配置能够提交多少 token。 if dynamic_batch_size == req_num * (verified_draft_step + 1): for depth in range(1, verified_draft_step + 1): survival = float(np.mean(accept_lengths > depth)) diff --git a/unit_tests/server/test_mtp_start_args.py b/unit_tests/server/test_mtp_start_args.py new file mode 100644 index 0000000000..4d48bd8182 --- /dev/null +++ b/unit_tests/server/test_mtp_start_args.py @@ -0,0 +1,12 @@ +import pytest + +from lightllm.server.api_start import _launch_subprocesses +from lightllm.server.core.objs.start_args_type import StartArgs + + +def test_mtp_requires_cuda_graph(monkeypatch): + monkeypatch.setattr("lightllm.server.api_start._set_envs_and_config", lambda args: None) + args = StartArgs(mtp_mode="vanilla_no_att", disable_cudagraph=True) + + with pytest.raises(AssertionError, match="--disable_cudagraph is not supported"): + _launch_subprocesses(args) From 4921565c62ff3783b39e7b3fa6476784a28237ca Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 00:33:52 +0000 Subject: [PATCH 043/103] refactor: add common MTP planner interface --- .../model_infer/mtp_speculative/engine.py | 5 ++- .../mtp_speculative/planner/__init__.py | 3 +- .../mtp_speculative/planner/base.py | 42 +++++++++++++++++++ .../mtp_speculative/planner/dspark.py | 8 +++- .../mtp_speculative/planner/fixed.py | 15 ++++++- .../mtp_speculative/planner/lightspec.py | 3 +- .../mtp_speculative/test_planner.py | 14 +++++-- 7 files changed, 79 insertions(+), 11 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 1301b2470e..41dc7236a0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -13,6 +13,7 @@ ) from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager from lightllm.server.router.model_infer.mtp_speculative.planner import ( + BaseMtpPlanner, DSparkPlanner, FixedSpecPlanner, LightSpecPlanner, @@ -39,7 +40,7 @@ def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp, ) - self.planner = self._build_decode_planner() + self.planner: BaseMtpPlanner = self._build_decode_planner() # Prefill draft-state initialization. @@ -296,7 +297,7 @@ def free_unused_decode_mem( # Construction helpers. - def _build_decode_planner(self): + def _build_decode_planner(self) -> BaseMtpPlanner: if not self.enable_dynmaic_mtp: return FixedSpecPlanner(max_draft_step=self.backend.max_draft_step) if self.spec_mode == "dspark": diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py index a48dcb7444..6db1b911be 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py @@ -1,10 +1,11 @@ -from lightllm.server.router.model_infer.mtp_speculative.planner.base import SpecDecodePlan +from lightllm.server.router.model_infer.mtp_speculative.planner.base import BaseMtpPlanner, SpecDecodePlan from lightllm.server.router.model_infer.mtp_speculative.planner.dspark import DSparkPlanner from lightllm.server.router.model_infer.mtp_speculative.planner.fixed import FixedSpecPlanner from lightllm.server.router.model_infer.mtp_speculative.planner.lightspec import LightSpecPlanner __all__ = [ + "BaseMtpPlanner", "DSparkPlanner", "FixedSpecPlanner", "LightSpecPlanner", diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py index 1e8f9cd91b..7fc6aaa33d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -1,5 +1,6 @@ from __future__ import annotations +from abc import ABC, abstractmethod from dataclasses import dataclass from typing import List, Optional @@ -41,6 +42,47 @@ def filter_reqs(self, reqs: List, selected_row_mask_cpu) -> List: return [req for req, selected in zip(reqs, selected_row_mask_cpu.tolist()) if selected] +class BaseMtpPlanner(ABC): + """定义 SpecEngine 与不同 MTP 规划器之间的统一调用接口。""" + + @abstractmethod + def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: + """为当前 decode 迭代生成执行计划。 + + Args: + decode_reqs: 当前参与 decode 的逻辑请求列表。规划器可以读取请求的 + 输出进度,判断请求是否已经持有上一轮生成的 draft proposal。 + original_batch_size: 进入动态压缩前的物理 verify 行数。 + + Returns: + 本轮 target verify 使用的动态 batch size、下一轮需要生成的 draft + step,以及描述当前 proposal 布局的上一轮 draft step。 + """ + + raise NotImplementedError + + @abstractmethod + def update_feedback( + self, + plan: SpecDecodePlan, + req_num: int, + accept_lengths, + schedule_scores=None, + ) -> None: + """在本轮 verify 完成后更新规划器的运行时统计。 + + Args: + plan: 本轮 decode 实际采用的执行计划,用于确定被验证的配置。 + req_num: 本轮逻辑请求数量。 + accept_lengths: 每个请求本轮提交的 token 数量,包含必然提交的 + target token。 + schedule_scores: proposer 可选提供的调度分数。LightSpec 主要使用 + accept_lengths,DSpark 使用置信度分数,固定规划器忽略所有反馈。 + """ + + raise NotImplementedError + + class _InferCostMsTable: def __init__(self) -> None: self.infer_cost_ms_table = SortedDict() diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index 9f55a14300..4b28484a2f 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -5,10 +5,14 @@ import numpy as np -from lightllm.server.router.model_infer.mtp_speculative.planner.base import SpecDecodePlan, _InferCostMsTable +from lightllm.server.router.model_infer.mtp_speculative.planner.base import ( + BaseMtpPlanner, + SpecDecodePlan, + _InferCostMsTable, +) -class DSparkPlanner: +class DSparkPlanner(BaseMtpPlanner): """DSpark's confidence-based verify-capacity planner. DSpark always drafts a complete block. Confidence from one proposal selects diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py index a3d76015ea..fe55b6d3a6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py @@ -2,10 +2,10 @@ from typing import List -from lightllm.server.router.model_infer.mtp_speculative.planner.base import SpecDecodePlan +from lightllm.server.router.model_infer.mtp_speculative.planner.base import BaseMtpPlanner, SpecDecodePlan -class FixedSpecPlanner: +class FixedSpecPlanner(BaseMtpPlanner): """Planner for fixed-width speculative decoding.""" def __init__(self, max_draft_step: int) -> None: @@ -17,3 +17,14 @@ def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: draft_step=self.max_draft_step, pre_draft_step=self.max_draft_step, ) + + def update_feedback( + self, + plan: SpecDecodePlan, + req_num: int, + accept_lengths, + schedule_scores=None, + ) -> None: + """固定规划不根据运行反馈调整 batch size 或 draft step。""" + + return diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index 1b82285a3d..beeb942c0d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -5,13 +5,14 @@ import numpy as np from lightllm.server.router.model_infer.mtp_speculative.planner.base import ( + BaseMtpPlanner, SpecDecodePlan, _EMAValue, _InferCostMsTable, ) -class LightSpecPlanner: +class LightSpecPlanner(BaseMtpPlanner): """Choose the current verify budget and the next draft configuration. LightSpec evaluates a runtime configuration as ``(N, B, d)``: logical diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 38eb3e0386..32d7c88ebe 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -6,6 +6,7 @@ from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine from lightllm.server.router.model_infer.mtp_speculative.planner import ( + BaseMtpPlanner, DSparkPlanner, FixedSpecPlanner, LightSpecPlanner, @@ -70,8 +71,10 @@ def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): def test_fixed_planner_returns_static_plan(): - plan = FixedSpecPlanner(max_draft_step=3).plan(decode_reqs=build_decode_reqs(4), original_batch_size=16) + planner = FixedSpecPlanner(max_draft_step=3) + plan = planner.plan(decode_reqs=build_decode_reqs(4), original_batch_size=16) + assert isinstance(planner, BaseMtpPlanner) assert not plan.is_dynamic assert plan.dynamic_batch_size is None assert plan.draft_step == plan.pre_draft_step == 3 @@ -87,11 +90,16 @@ def test_infer_cost_candidates_include_feasible_boundaries(): def test_engine_routes_only_dspark_to_the_confidence_planner(): - assert isinstance(build_planner("eagle3", enable_dynmaic_mtp=False), FixedSpecPlanner) - assert isinstance(build_planner("dspark"), DSparkPlanner) + fixed_planner = build_planner("eagle3", enable_dynmaic_mtp=False) + dspark_planner = build_planner("dspark") + assert isinstance(fixed_planner, FixedSpecPlanner) + assert isinstance(dspark_planner, DSparkPlanner) + assert isinstance(fixed_planner, BaseMtpPlanner) + assert isinstance(dspark_planner, BaseMtpPlanner) vanilla_planner = build_planner("vanilla_with_att") assert isinstance(vanilla_planner, LightSpecPlanner) + assert isinstance(vanilla_planner, BaseMtpPlanner) assert vanilla_planner.draft_steps == (1, 2, 3) vanilla_no_att_planner = build_planner("vanilla_no_att") From 46b04e87a1bb380ebd81baf1cf39e4464af7405b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 00:35:56 +0000 Subject: [PATCH 044/103] chore: add DP MTP extension packages --- .../router/model_infer/mtp_speculative/dp_planner/__init__.py | 1 + .../router/model_infer/mtp_speculative/dp_proposers/__init__.py | 1 + 2 files changed, 2 insertions(+) create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py new file mode 100644 index 0000000000..0a00c771a7 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py @@ -0,0 +1 @@ +"""DP-specific MTP planner implementations.""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py new file mode 100644 index 0000000000..8274fbb5a4 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py @@ -0,0 +1 @@ +"""DP-specific MTP proposer implementations.""" From 45a21d4c332aa36d6b60b68749b01c8dc723ebce Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 01:58:48 +0000 Subject: [PATCH 045/103] refactor: separate DP MTP engines and proposers --- .../model_infer/mode_backend/base_backend.py | 16 +- .../mode_backend/chunked_prefill/impl.py | 9 + .../mode_backend/dp_backend/impl.py | 104 ++-- .../model_infer/mtp_speculative/__init__.py | 4 +- .../model_infer/mtp_speculative/dp_engine.py | 125 +++++ .../mtp_speculative/dp_overlap_engine.py | 72 +++ .../dp_overlap_planner/__init__.py | 9 + .../dp_overlap_planner/base.py | 11 + .../dp_overlap_planner/fixed.py | 11 + .../dp_overlap_proposers/__init__.py | 47 ++ .../dp_overlap_proposers/base.py | 43 ++ .../dp_overlap_proposers/eagle3.py | 97 ++++ .../dp_overlap_proposers/eagle_no_att.py | 92 ++++ .../dp_overlap_proposers/eagle_utils.py | 287 ++++++++++++ .../dp_overlap_proposers/eagle_with_att.py | 92 ++++ .../dp_overlap_proposers/vanilla_no_att.py | 82 ++++ .../dp_overlap_proposers/vanilla_utils.py | 86 ++++ .../dp_overlap_proposers/vanilla_with_att.py | 82 ++++ .../mtp_speculative/dp_planner/__init__.py | 10 +- .../mtp_speculative/dp_planner/base.py | 11 + .../mtp_speculative/dp_planner/fixed.py | 11 + .../mtp_speculative/dp_proposers/__init__.py | 39 +- .../mtp_speculative/dp_proposers/base.py | 7 + .../mtp_speculative/dp_proposers/eagle3.py | 44 ++ .../dp_proposers/eagle_no_att.py | 41 ++ .../dp_proposers/eagle_with_att.py | 41 ++ .../dp_proposers/vanilla_no_att.py | 32 ++ .../dp_proposers/vanilla_with_att.py | 32 ++ .../model_infer/mtp_speculative/engine.py | 53 +-- .../mtp_speculative/proposers/__init__.py | 22 +- .../mtp_speculative/proposers/base.py | 36 +- .../mtp_speculative/proposers/dflash.py | 29 +- .../mtp_speculative/proposers/dspark.py | 29 +- .../mtp_speculative/proposers/eagle3.py | 53 ++- .../mtp_speculative/proposers/eagle_mtp.py | 443 ------------------ .../mtp_speculative/proposers/eagle_no_att.py | 40 ++ .../mtp_speculative/proposers/eagle_utils.py | 186 ++++++++ .../proposers/eagle_with_att.py | 40 ++ .../proposers/parallel_block.py | 100 ---- .../proposers/parallel_block_utils.py | 97 ++++ .../mtp_speculative/proposers/vanilla_mtp.py | 113 ----- .../proposers/vanilla_no_att.py | 42 ++ .../proposers/vanilla_utils.py | 75 +++ .../proposers/vanilla_with_att.py | 42 ++ .../test_dp_overlap_spec_engine.py | 234 +++++++++ .../mode_backend/test_dp_spec_engine.py | 118 ----- .../mtp_speculative/test_eagle_overlap.py | 15 +- .../mtp_speculative/test_planner.py | 159 ++++++- .../mtp_speculative/test_vanilla_overlap.py | 65 +++ 49 files changed, 2550 insertions(+), 978 deletions(-) create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_engine.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/base.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/fixed.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_planner/base.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_planner/fixed.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/base.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py create mode 100644 unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py delete mode 100644 unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py create mode 100644 unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py 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 1a0e867d4e..e0d1b3924a 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -36,7 +36,6 @@ get_radix_tree_merge_update_delta, ) from lightllm.distributed import dist_group_manager -from lightllm.server.router.model_infer.mtp_speculative import SpecEngine from lightllm.distributed.communication_op import ( all_gather_into_tensor, all_reduce, @@ -243,7 +242,9 @@ def init_model(self, kvargs): # 只会在 pd pd 模式下才会使用,用于上传分块传输任务是否成功。 self.shm_pd_trans_io_buffer = ShmObjsIOBuffer(tail_str="pd") - self.init_spec_engine(model_kvargs) + if self.args.mtp_mode is not None: + self.init_mtp_draft_model(model_kvargs) + self.init_spec_engine() if self.args.enable_cpu_cache: self.multi_level_cache_module = MultiLevelKvCacheModule(self) @@ -296,17 +297,6 @@ def prefill(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): def decode(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): raise NotImplementedError() - def init_spec_engine(self, main_kvargs: dict): - if self.args.mtp_mode is None: - return - self.init_mtp_draft_model(main_kvargs) - self.spec_engine = SpecEngine( - backend=self, - spec_mode=self.args.mtp_mode, - enable_dynmaic_mtp=self.args.mtp_dynamic_verify, - ) - return - def init_mtp_draft_model(self, main_kvargs: dict): self.max_draft_step = self.args.mtp_step self.draft_models = [] diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 8661208b18..928b97bdc6 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -12,6 +12,7 @@ from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample from lightllm.server.router.model_infer.infer_batch import g_infer_context from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_current_device_id from .control_state import ControlState @@ -40,6 +41,14 @@ def __init__(self) -> None: self.classed_req_strict_prefill = False return + def init_spec_engine(self): + self.spec_engine = SpecEngine( + backend=self, + spec_mode=self.args.mtp_mode, + enable_dynmaic_mtp=self.args.mtp_dynamic_verify, + ) + return + def infer_loop(self): torch.cuda.set_device(get_current_device_id()) try: diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index fed67950e0..e7ef122549 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -15,6 +15,8 @@ from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.envs_utils import get_env_start_args from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from .control_state import DPControlState @@ -65,6 +67,22 @@ def __init__(self) -> None: self.classed_req_strict_prefill = False return + def init_spec_engine(self): + engine_kwargs = dict( + backend=self, + spec_mode=self.args.mtp_mode, + enable_dynmaic_mtp=self.args.mtp_dynamic_verify, + ) + self.spec_engine = DPSpecEngine(**engine_kwargs) + self.dp_overlap_spec_engine = DPOverlapSpecEngine(**engine_kwargs) + self.prefill_draft_engine = ( + self.dp_overlap_spec_engine if self.enable_prefill_microbatch_overlap else self.spec_engine + ) + self.decode_draft_engine = ( + self.dp_overlap_spec_engine if self.enable_decode_microbatch_overlap else self.spec_engine + ) + return + @staticmethod def _build_padded_next_token_ids( token_ids: torch.Tensor | None, @@ -458,7 +476,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] copy_len=req_num, device=model_input.b_req_idx.device, ) - self.spec_engine.build_draft_state_from_prefill( + self.prefill_draft_engine.build_draft_state_from_prefill( target_model_input=model_input, target_model_output=model_output, next_token_ids=draft_next_token_ids_gpu, @@ -611,38 +629,27 @@ def _draft_decode_vanilla( mtp_accept_len: torch.Tensor, req_num: int, ): - all_next_token_ids = [] - # share some inference info with the main model - draft_model_input = model_input - draft_hidden = model_output.mtp_collector.spec_hidden - draft_next_token_ids_gpu = self._build_padded_next_token_ids( + padded_next_token_ids = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input.batch_size, copy_len=req_num, device=model_input.b_req_idx.device, ) - - all_next_token_ids.append(draft_next_token_ids_gpu) - - # process the draft model output - for draft_model_idx in range(self.max_draft_step): - draft_model_input.input_ids = draft_next_token_ids_gpu - draft_model_input.mtp_draft_input_hiddens = draft_hidden - # spec decode: MTP - draft_model_output: ModelOutput = self.draft_models[draft_model_idx].forward(draft_model_input) - draft_hidden = draft_model_output.mtp_collector.spec_hidden - draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) - all_next_token_ids.append(draft_next_token_ids_gpu) + proposal = self.decode_draft_engine.propose_next( + main_model_input=model_input, + main_model_output=model_output, + next_token_ids=padded_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + ) if req_num > 0: - stacked_next_token_ids = torch.stack(all_next_token_ids, dim=1)[:req_num] self.spec_engine.scatter_next_tokens( - all_next_token_ids=stacked_next_token_ids, + all_next_token_ids=proposal.token_ids[:req_num], b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx[:req_num], mtp_accept_len=mtp_accept_len, ) - return None + return proposal.extra_mem_indexes_cpu def _draft_decode_eagle( self, @@ -682,12 +689,11 @@ def _draft_decode_eagle( # agreement. The proposer still follows the common topology: one # full-row extend, followed by autoregressive drafting over one row per # (real or HOLD) request. - proposal = self.spec_engine.propose_next( + proposal = self.decode_draft_engine.propose_next( main_model_input=model_input, main_model_output=model_output, next_token_ids=padded_next_token_ids, b_req_mtp_start_loc=padded_start_locs, - draft_step=self.max_draft_step, accept_len=padded_accept_len, ) @@ -759,7 +765,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I source_start=req_num0, ) - self.spec_engine.build_draft_state_from_prefill_overlap( + self.prefill_draft_engine.build_draft_state_from_prefill_overlap( target_model_input0=model_input0, target_model_output0=model_output0, next_token_ids0=draft_next_token_ids_gpu0, @@ -942,21 +948,14 @@ def _draft_decode_vanilla_overlap( req_num0: int = 0, req_num1: int = 0, ): - all_next_token_ids = [] - all_next_token_ids.append(next_token_ids) - # share some inference info with the main model - draft_model_input0, draft_model_input1 = model_input0, model_input1 - draft_hidden0 = model_output0.mtp_collector.spec_hidden - draft_hidden1 = model_output1.mtp_collector.spec_hidden - - draft_next_token_ids_gpu0 = self._build_padded_next_token_ids( + padded_next_token_ids0 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input0.batch_size, copy_len=req_num0, device=model_input0.b_req_idx.device, source_start=0, ) - draft_next_token_ids_gpu1 = self._build_padded_next_token_ids( + padded_next_token_ids1 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input1.batch_size, copy_len=req_num1, @@ -964,35 +963,27 @@ def _draft_decode_vanilla_overlap( source_start=req_num0, ) - # process the draft model output - for draft_model_idx in range(self.max_draft_step): - draft_model_input0.input_ids = draft_next_token_ids_gpu0 - draft_model_input0.mtp_draft_input_hiddens = draft_hidden0 - draft_model_input1.input_ids = draft_next_token_ids_gpu1 - draft_model_input1.mtp_draft_input_hiddens = draft_hidden1 - - draft_model_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( - draft_model_input0, draft_model_input1 - ) - draft_hidden0 = draft_model_output0.mtp_collector.spec_hidden - draft_hidden1 = draft_model_output1.mtp_collector.spec_hidden - - draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) - draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) - draft_next_token_ids = torch.cat( - (draft_next_token_ids_gpu0[0:req_num0], draft_next_token_ids_gpu1[0:req_num1]), dim=0 - ) - all_next_token_ids.append(draft_next_token_ids) + proposal = self.decode_draft_engine.propose_next_overlap( + main_model_input0=model_input0, + main_model_output0=model_output0, + next_token_ids0=padded_next_token_ids0, + real_verify_rows0=req_num0, + accept_len0=None, + main_model_input1=model_input1, + main_model_output1=model_output1, + next_token_ids1=padded_next_token_ids1, + real_verify_rows1=req_num1, + accept_len1=None, + ) if req_num0 + req_num1 > 0: - stacked_next_token_ids = torch.stack(all_next_token_ids, dim=1) self.spec_engine.scatter_next_tokens( - all_next_token_ids=stacked_next_token_ids, + all_next_token_ids=proposal.token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, ) - return None + return proposal.extra_mem_indexes_cpu def _draft_decode_eagle_overlap( self, @@ -1044,7 +1035,7 @@ def _draft_decode_eagle_overlap( mtp_accept_len[real_request_num0 : real_request_num0 + real_request_num1] ) - proposal = self.spec_engine.propose_next_overlap( + proposal = self.decode_draft_engine.propose_next_overlap( main_model_input0=model_input0, main_model_output0=model_output0, next_token_ids0=padded_next_token_ids0, @@ -1055,7 +1046,6 @@ def _draft_decode_eagle_overlap( next_token_ids1=padded_next_token_ids1, real_verify_rows1=req_num1, accept_len1=padded_accept_len1, - draft_step=self.max_draft_step, ) if req_num0 + req_num1 > 0: diff --git a/lightllm/server/router/model_infer/mtp_speculative/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/__init__.py index 779438e356..4944600fbd 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/__init__.py @@ -1,4 +1,6 @@ +from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine -__all__ = ["SpecEngine"] +__all__ = ["DPSpecEngine", "DPOverlapSpecEngine", "SpecEngine"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py new file mode 100644 index 0000000000..2d5dce8eee --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py @@ -0,0 +1,125 @@ +from typing import List, Optional, Tuple + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.mtp_utils import ( + linear_att_mtp_state_index_update, + mtp_scatter_next_token_ids, + mtp_verify, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_planner import BaseDpPlanner, build_dp_planner +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers import build_dp_spec_proposer +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + + +class DPSpecEngine: + """普通 DP prefill/decode 使用的 MTP engine。""" + + def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: + self.backend = backend + self.proposer: BaseDpProposer = build_dp_spec_proposer( + spec_mode=spec_mode, + backend=backend, + enable_dynmaic_mtp=enable_dynmaic_mtp, + ) + self.planner: BaseDpPlanner = build_dp_planner(backend=backend) + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + self.proposer.build_draft_state_from_prefill( + target_model_input=target_model_input, + target_model_output=target_model_output, + next_token_ids=next_token_ids, + ) + + def verify_tokens( + self, + next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + b_mtp_index: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + accept_lengths, accepted_index = mtp_verify( + req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + new_next_token_ids=next_token_ids, + b_req_idx=b_req_idx, + ) + if self.backend.is_linear_att_mixed_model: + assert b_mtp_index is not None + linear_att_mtp_state_index_update( + req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_req_idx=b_req_idx, + b_mtp_index=b_mtp_index, + accepted_index=accepted_index, + verify_width=self.backend.max_draft_step + 1, + ) + return accept_lengths, accepted_index + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + accept_len: Optional[torch.Tensor] = None, + ) -> SpecProposal: + return self.proposer.propose_next( + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=self.planner.get_draft_step(), + accept_len=accept_len, + ) + + def scatter_next_tokens( + self, + b_req_mtp_start_loc: torch.Tensor, + all_next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + mtp_accept_len: torch.Tensor, + schedule_scores: Optional[torch.Tensor] = None, + ) -> None: + mtp_scatter_next_token_ids( + req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=b_req_idx, + mtp_accept_len=mtp_accept_len, + req_to_next_token_scores=( + self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_scores + if schedule_scores is not None + else None + ), + schedule_scores=schedule_scores, + ) + + def record_request_spec_metrics( + self, + decode_reqs: List, + accept_lengths_cpu: torch.Tensor, + ) -> None: + """累计 DP 固定 verify layout 下每个请求的 MTP 指标。""" + + if not self.backend.is_master_in_dp: + return + + accept_lengths = accept_lengths_cpu.tolist() + assert len(accept_lengths) == len(decode_reqs) + for req, accept_len in zip(decode_reqs, accept_lengths): + req.update_mtp_accepted_token_num(accept_token_num=accept_len - 1) + verify_token_num = req.mtp_step + 1 + if verify_token_num > 0: + req.update_mtp_verify_token_num(verify_token_num=verify_token_num) + req.update_mtp_verify_step_num(verify_step_num=1) + + +__all__ = ["DPSpecEngine"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py new file mode 100644 index 0000000000..ca7ca8a11f --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -0,0 +1,72 @@ +from typing import Optional + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner import ( + BaseDpOverlapPlanner, + build_dp_overlap_planner, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers import build_dp_overlap_spec_proposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + + +class DPOverlapSpecEngine: + """双 microbatch overlap draft 流程使用的 DP MTP engine。""" + + def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: + self.proposer: BaseDpOverlapProposer = build_dp_overlap_spec_proposer( + spec_mode=spec_mode, + backend=backend, + enable_dynmaic_mtp=enable_dynmaic_mtp, + ) + self.planner: BaseDpOverlapPlanner = build_dp_overlap_planner(backend=backend) + + def build_draft_state_from_prefill_overlap( + self, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + ) -> None: + self.proposer.build_draft_state_from_prefill_overlap( + target_model_input0=target_model_input0, + target_model_output0=target_model_output0, + next_token_ids0=next_token_ids0, + target_model_input1=target_model_input1, + target_model_output1=target_model_output1, + next_token_ids1=next_token_ids1, + ) + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: Optional[torch.Tensor], + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: Optional[torch.Tensor], + ) -> SpecProposal: + return self.proposer.propose_next_overlap( + main_model_input0=main_model_input0, + main_model_output0=main_model_output0, + next_token_ids0=next_token_ids0, + real_verify_rows0=real_verify_rows0, + accept_len0=accept_len0, + main_model_input1=main_model_input1, + main_model_output1=main_model_output1, + next_token_ids1=next_token_ids1, + real_verify_rows1=real_verify_rows1, + accept_len1=accept_len1, + draft_step=self.planner.get_draft_step(), + ) + + +__all__ = ["DPOverlapSpecEngine"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py new file mode 100644 index 0000000000..aee2c5b8d3 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py @@ -0,0 +1,9 @@ +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner.base import BaseDpOverlapPlanner +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner.fixed import FixedDpOverlapPlanner + + +def build_dp_overlap_planner(*, backend) -> BaseDpOverlapPlanner: + return FixedDpOverlapPlanner(draft_step=backend.max_draft_step) + + +__all__ = ["BaseDpOverlapPlanner", "FixedDpOverlapPlanner", "build_dp_overlap_planner"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/base.py new file mode 100644 index 0000000000..a809a28c6e --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/base.py @@ -0,0 +1,11 @@ +from abc import ABC, abstractmethod + + +class BaseDpOverlapPlanner(ABC): + """DP overlap draft 配置的基础规划接口。""" + + @abstractmethod + def get_draft_step(self) -> int: + """返回当前 overlap proposal 使用的 draft step。""" + + raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/fixed.py new file mode 100644 index 0000000000..3c96f9a0fc --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/fixed.py @@ -0,0 +1,11 @@ +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner.base import BaseDpOverlapPlanner + + +class FixedDpOverlapPlanner(BaseDpOverlapPlanner): + """为 DP overlap 固定 shape decode 返回启动时配置的 draft step。""" + + def __init__(self, draft_step: int) -> None: + self.draft_step = int(draft_step) + + def get_draft_step(self) -> int: + return self.draft_step diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py new file mode 100644 index 0000000000..63841bcaf3 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py @@ -0,0 +1,47 @@ +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer + + +def build_dp_overlap_spec_proposer( + *, + spec_mode: str, + backend, + enable_dynmaic_mtp: bool, +) -> "BaseDpOverlapProposer": + if spec_mode == "vanilla_with_att": + from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_with_att import ( + DpOverlapVanillaWithAttProposer, + ) + + return DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "vanilla_no_att": + from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_no_att import ( + DpOverlapVanillaNoAttProposer, + ) + + return DpOverlapVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "eagle_with_att": + from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_with_att import ( + DpOverlapEagleWithAttProposer, + ) + + return DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "eagle_no_att": + from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_no_att import ( + DpOverlapEagleNoAttProposer, + ) + + return DpOverlapEagleNoAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "eagle3": + from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import ( + DpOverlapEagle3Proposer, + ) + + return DpOverlapEagle3Proposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + + raise ValueError(f"unsupported DP overlap speculative mode: {spec_mode}") + + +__all__ = ["build_dp_overlap_spec_proposer"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py new file mode 100644 index 0000000000..fe25b40b9d --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -0,0 +1,43 @@ +from abc import ABC, abstractmethod + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal + + +class BaseDpOverlapProposer(BaseSpecProposer, ABC): + """DP proposer 的完整接口,扩展双 microbatch overlap 操作。""" + + @abstractmethod + def build_draft_state_from_prefill_overlap( + self, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + ) -> None: + """Build draft state from two overlapped target-prefill microbatches.""" + + raise NotImplementedError + + @abstractmethod + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor | None, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor | None, + draft_step: int, + ) -> SpecProposal: + """Generate one proposal from two DP-overlapped decode microbatches.""" + + raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py new file mode 100644 index 0000000000..e3ddd2a3be --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -0,0 +1,97 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( + build_dp_eagle_draft_state_from_prefill_overlap, + propose_next_dp_eagle_autoregressive_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + propose_next_eagle, +) + + +class DpOverlapEagle3Proposer(BaseDpOverlapProposer): + """DP ``eagle3`` proposer。""" + + def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: + return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def build_draft_state_from_prefill_overlap( + self, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + ) -> None: + build_dp_eagle_draft_state_from_prefill_overlap( + self, + target_model_input0, + target_model_output0, + next_token_ids0, + target_model_input1, + target_model_output1, + next_token_ids1, + ) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + self._map_draft_token_ids, + ) + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor | None, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor | None, + draft_step: int, + ) -> SpecProposal: + return propose_next_dp_eagle_autoregressive_overlap( + self, + main_model_input0, + main_model_output0, + next_token_ids0, + real_verify_rows0, + accept_len0, + main_model_input1, + main_model_output1, + next_token_ids1, + real_verify_rows1, + accept_len1, + draft_step, + self._map_draft_token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py new file mode 100644 index 0000000000..7e147a0106 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -0,0 +1,92 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( + build_dp_eagle_draft_state_from_prefill_overlap, + propose_next_dp_eagle_fixed_layout_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + propose_next_eagle, +) + + +class DpOverlapEagleNoAttProposer(BaseDpOverlapProposer): + """DP ``eagle_no_att`` proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def build_draft_state_from_prefill_overlap( + self, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + ) -> None: + build_dp_eagle_draft_state_from_prefill_overlap( + self, + target_model_input0, + target_model_output0, + next_token_ids0, + target_model_input1, + target_model_output1, + next_token_ids1, + ) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + lambda token_ids: token_ids, + ) + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor | None, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor | None, + draft_step: int, + ) -> SpecProposal: + return propose_next_dp_eagle_fixed_layout_overlap( + self, + main_model_input0, + main_model_output0, + next_token_ids0, + real_verify_rows0, + main_model_input1, + main_model_output1, + next_token_ids1, + real_verify_rows1, + draft_step, + lambda token_ids: token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py new file mode 100644 index 0000000000..bbaa1b5202 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -0,0 +1,287 @@ +"""DP overlap EAGLE proposer 共享辅助函数。""" + +from __future__ import annotations + +from typing import Callable + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + generate_eagle_token_ids, + prepare_eagle_verify_extend_input, +) + + +def build_dp_eagle_draft_state_from_prefill_overlap( + proposer: BaseDpOverlapProposer, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, +) -> None: + """使用两个 target prefill microbatch 初始化 EAGLE draft state。""" + + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + prepare_mtp_prefill_inputs( + model_input=target_model_input0, + b_next_token_ids=next_token_ids0, + mtp_draft_input_hiddens=target_model_output0.mtp_collector.spec_hidden, + ) + prepare_mtp_prefill_inputs( + model_input=target_model_input1, + b_next_token_ids=next_token_ids1, + mtp_draft_input_hiddens=target_model_output1.mtp_collector.spec_hidden, + ) + proposer.backend.draft_models[0].microbatch_overlap_prefill(target_model_input0, target_model_input1) + + +def pad_dp_step_mem_indexes( + real_mem_indexes: torch.Tensor, + request_capacity: int, + hold_mem_index: int, +) -> torch.Tensor: + padded = torch.full( + (request_capacity,), + hold_mem_index, + dtype=real_mem_indexes.dtype, + device=real_mem_indexes.device, + ) + padded[: real_mem_indexes.shape[0]].copy_(real_mem_indexes) + return padded + + +def propose_next_dp_eagle_autoregressive_overlap( + proposer: BaseDpOverlapProposer, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor | None, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor | None, + draft_step: int, + map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], +) -> SpecProposal: + """运行 DP EAGLE extend 后接单 token overlap decode 的 proposal 流程。""" + + verify_width = proposer.backend.max_draft_step + 1 + model_inputs = (main_model_input0, main_model_input1) + model_outputs = (main_model_output0, main_model_output1) + next_token_ids_by_batch = (next_token_ids0, next_token_ids1) + real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) + accept_lens_by_batch = (accept_len0, accept_len1) + position_deltas_by_batch = tuple(model_input.b_position_delta for model_input in model_inputs) + max_kv_seq_lens_by_batch = tuple(model_input.max_kv_seq_len for model_input in model_inputs) + + request_capacities_by_batch = [] + real_request_counts = [] + accepted_tail_rows_by_batch = [] + for model_input, model_output, token_ids, real_verify_row_count, accept_len in zip( + model_inputs, + model_outputs, + next_token_ids_by_batch, + real_verify_row_counts, + accept_lens_by_batch, + ): + request_capacity = model_input.batch_size // verify_width + real_request_count = real_verify_row_count // verify_width + starts = torch.arange( + 0, + model_input.batch_size, + verify_width, + dtype=torch.int32, + device=token_ids.device, + ) + accepted_tail_rows = (starts + accept_len - 1).long() + request_capacities_by_batch.append(request_capacity) + real_request_counts.append(real_request_count) + accepted_tail_rows_by_batch.append(accepted_tail_rows) + prepare_eagle_verify_extend_input( + model_input=model_input, + input_ids=token_ids, + target_hidden=model_output.mtp_collector.spec_hidden, + ) + + verify_row_count = real_verify_rows0 + real_verify_rows1 + proposal_token_ids = next_token_ids0.new_full( + (verify_row_count, draft_step + 1), + fill_value=1, + ) + proposal_token_ids[:real_verify_rows0, 0].copy_(next_token_ids0[:real_verify_rows0]) + proposal_token_ids[real_verify_rows0:, 0].copy_(next_token_ids1[:real_verify_rows1]) + + draft_model = proposer.backend.draft_models[0] + extend_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) + + draft_token_ids_by_batch = [] + draft_hiddens_by_batch = [] + draft_seq_lens_by_batch = [] + draft_req_indices_by_batch = [] + proposal_rows_by_batch = [] + proposal_row_offsets = (0, real_verify_rows0) + for batch_index, (model_input, extend_output, accepted_tail_rows, real_request_count) in enumerate( + zip(model_inputs, extend_outputs, accepted_tail_rows_by_batch, real_request_counts) + ): + accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) + draft_token_ids = generate_eagle_token_ids(proposer, accepted_tail_output, map_draft_token_ids) + draft_token_ids_by_batch.append(draft_token_ids) + draft_hiddens_by_batch.append(extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows)) + draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) + draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) + proposal_rows = accepted_tail_rows[:real_request_count] + proposal_row_offsets[batch_index] + proposal_rows_by_batch.append(proposal_rows) + proposal_token_ids[proposal_rows, 1] = draft_token_ids[:real_request_count] + + if draft_step == 1: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + ) + + for batch_index, model_input in enumerate(model_inputs): + model_input.is_prefill = False + model_input.batch_size = request_capacities_by_batch[batch_index] + model_input.b_req_idx = draft_req_indices_by_batch[batch_index] + model_input.b_mtp_index = torch.zeros_like(model_input.b_req_idx) + model_input.b_seq_len = draft_seq_lens_by_batch[batch_index] + model_input.b_position_delta = ( + position_deltas_by_batch[batch_index].index_select(0, accepted_tail_rows_by_batch[batch_index]) + if position_deltas_by_batch[batch_index] is not None + else None + ) + model_input.b_mark_shared_group = torch.ones_like(model_input.b_req_idx) + model_input.b_shared_seq_len = None + if len(model_input.multimodal_params) != model_input.batch_size: + empty_multimodal_params = {"images": [], "audios": []} + model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size + + total_real_request_count = sum(real_request_counts) + extra_mem_indexes_cpu = proposer.alloc_extra_mem_indexes(total_real_request_count * (draft_step - 1)) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) + hold_mem_index = proposer.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX + + for step in range(1, draft_step): + mem_start = (step - 1) * total_real_request_count + step_mem_indexes = extra_mem_indexes[mem_start : mem_start + total_real_request_count] + real_mem_start = 0 + for batch_index, model_input in enumerate(model_inputs): + real_request_count = real_request_counts[batch_index] + real_mem_indexes = step_mem_indexes[real_mem_start : real_mem_start + real_request_count] + real_mem_start += real_request_count + model_input.input_ids = draft_token_ids_by_batch[batch_index] + model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] + model_input.mem_indexes = pad_dp_step_mem_indexes( + real_mem_indexes=real_mem_indexes, + request_capacity=request_capacities_by_batch[batch_index], + hold_mem_index=hold_mem_index, + ) + model_input.max_kv_seq_len = max_kv_seq_lens_by_batch[batch_index] + step + model_input.total_token_num = model_input.batch_size * model_input.max_kv_seq_len + + draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + draft_token_ids = generate_eagle_token_ids(proposer, draft_output, map_draft_token_ids) + draft_token_ids_by_batch[batch_index] = draft_token_ids + draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden + draft_seq_lens_by_batch[batch_index].add_(1) + real_request_count = real_request_counts[batch_index] + proposal_token_ids[proposal_rows_by_batch[batch_index], step + 1] = draft_token_ids[:real_request_count] + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, + ) + + +def propose_next_dp_eagle_fixed_layout_overlap( + proposer: BaseDpOverlapProposer, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + draft_step: int, + map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], +) -> SpecProposal: + """保持 expanded verify-row layout 运行 DP EAGLE overlap decode。""" + + verify_width = proposer.backend.max_draft_step + 1 + model_inputs = (main_model_input0, main_model_input1) + real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) + real_request_counts = tuple(row_count // verify_width for row_count in real_verify_row_counts) + request_capacities_by_batch = tuple(model_input.batch_size // verify_width for model_input in model_inputs) + total_real_request_count = sum(real_request_counts) + + proposal_token_ids = next_token_ids0.new_empty((sum(real_verify_row_counts), draft_step + 1)) + proposal_token_ids[:real_verify_rows0, 0] = next_token_ids0[:real_verify_rows0] + proposal_token_ids[real_verify_rows0:, 0] = next_token_ids1[:real_verify_rows1] + + extra_mem_indexes_cpu = proposer.alloc_extra_mem_indexes(total_real_request_count * draft_step) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) + split = real_request_counts[0] * draft_step + extra_mem_indexes_by_batch = ( + extra_mem_indexes[:split], + extra_mem_indexes[split:], + ) + + draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] + draft_hiddens_by_batch = [ + main_model_output0.mtp_collector.spec_hidden, + main_model_output1.mtp_collector.spec_hidden, + ] + draft_model = proposer.backend.draft_models[0] + hold_mem_index = proposer.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX + + for step in range(draft_step): + for batch_index, model_input in enumerate(model_inputs): + model_input.input_ids = draft_token_ids_by_batch[batch_index] + model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] + + draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) + + for batch_index, (model_input, draft_output) in enumerate(zip(model_inputs, draft_outputs)): + model_input.b_seq_len += 1 + model_input.max_kv_seq_len += 1 + + real_request_count = real_request_counts[batch_index] + mem_start = step * real_request_count + step_mem_indexes = extra_mem_indexes_by_batch[batch_index][mem_start : mem_start + real_request_count] + step_mem_indexes = pad_dp_step_mem_indexes( + real_mem_indexes=step_mem_indexes, + request_capacity=request_capacities_by_batch[batch_index], + hold_mem_index=hold_mem_index, + ) + model_input.mem_indexes = torch.cat( + [ + model_input.mem_indexes.view(-1, verify_width)[:, 1:], + step_mem_indexes.view(-1, 1), + ], + dim=1, + ).view(-1) + + draft_token_ids_by_batch[batch_index] = generate_eagle_token_ids( + proposer, + draft_output, + map_draft_token_ids, + ) + draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden + + proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids_by_batch[0][:real_verify_rows0] + proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids_by_batch[1][:real_verify_rows1] + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py new file mode 100644 index 0000000000..d3e51c78f5 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -0,0 +1,92 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( + build_dp_eagle_draft_state_from_prefill_overlap, + propose_next_dp_eagle_fixed_layout_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + propose_next_eagle, +) + + +class DpOverlapEagleWithAttProposer(BaseDpOverlapProposer): + """DP ``eagle_with_att`` proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def build_draft_state_from_prefill_overlap( + self, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + ) -> None: + build_dp_eagle_draft_state_from_prefill_overlap( + self, + target_model_input0, + target_model_output0, + next_token_ids0, + target_model_input1, + target_model_output1, + next_token_ids1, + ) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + lambda token_ids: token_ids, + ) + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor | None, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor | None, + draft_step: int, + ) -> SpecProposal: + return propose_next_dp_eagle_fixed_layout_overlap( + self, + main_model_input0, + main_model_output0, + next_token_ids0, + real_verify_rows0, + main_model_input1, + main_model_output1, + next_token_ids1, + real_verify_rows1, + draft_step, + lambda token_ids: token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py new file mode 100644 index 0000000000..6fa226a5e4 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -0,0 +1,82 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( + build_dp_chained_mtp_draft_state_from_prefill_overlap, + propose_next_dp_chained_mtp_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + build_chained_mtp_draft_state_from_prefill, + propose_next_chained_mtp, +) + + +class DpOverlapVanillaNoAttProposer(BaseDpOverlapProposer): + """DP ``vanilla_no_att`` proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_chained_mtp_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def build_draft_state_from_prefill_overlap( + self, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + ) -> None: + build_dp_chained_mtp_draft_state_from_prefill_overlap( + self, + target_model_input0, + target_model_output0, + next_token_ids0, + target_model_input1, + target_model_output1, + next_token_ids1, + ) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor | None, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor | None, + draft_step: int, + ) -> SpecProposal: + return propose_next_dp_chained_mtp_overlap( + self, + main_model_input0, + main_model_output0, + next_token_ids0, + real_verify_rows0, + main_model_input1, + main_model_output1, + next_token_ids1, + real_verify_rows1, + draft_step, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py new file mode 100644 index 0000000000..b60066fd1b --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py @@ -0,0 +1,86 @@ +"""DP overlap Vanilla proposer 共享辅助函数。""" + +from __future__ import annotations + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + + +def build_dp_chained_mtp_draft_state_from_prefill_overlap( + proposer: BaseDpOverlapProposer, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, +) -> None: + """为两个 DP microbatch 依次构建 Vanilla chained draft state。""" + + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + model_inputs = (target_model_input0, target_model_input1) + draft_hiddens = [ + target_model_output0.mtp_collector.spec_hidden, + target_model_output1.mtp_collector.spec_hidden, + ] + draft_token_ids = [next_token_ids0, next_token_ids1] + + for draft_model in proposer.backend.draft_models: + for batch_index, model_input in enumerate(model_inputs): + prepare_mtp_prefill_inputs( + model_input=model_input, + b_next_token_ids=draft_token_ids[batch_index], + mtp_draft_input_hiddens=draft_hiddens[batch_index], + ) + draft_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden + draft_token_ids[batch_index] = proposer.backend._gen_argmax_token_ids(draft_output) + + +def propose_next_dp_chained_mtp_overlap( + proposer: BaseDpOverlapProposer, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + draft_step: int, +) -> SpecProposal: + """为两个 DP microbatch 运行 Vanilla chained overlap decode。""" + + model_inputs = (main_model_input0, main_model_input1) + real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) + draft_token_ids = [next_token_ids0, next_token_ids1] + draft_hiddens = [ + main_model_output0.mtp_collector.spec_hidden, + main_model_output1.mtp_collector.spec_hidden, + ] + proposal_token_ids = next_token_ids0.new_empty((sum(real_verify_rows), draft_step + 1)) + proposal_token_ids[:real_verify_rows0, 0] = next_token_ids0[:real_verify_rows0] + proposal_token_ids[real_verify_rows0:, 0] = next_token_ids1[:real_verify_rows1] + + for step in range(draft_step): + for batch_index, model_input in enumerate(model_inputs): + model_input.input_ids = draft_token_ids[batch_index] + model_input.mtp_draft_input_hiddens = draft_hiddens[batch_index] + + draft_outputs = proposer.backend.draft_models[step].microbatch_overlap_decode(*model_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden + draft_token_ids[batch_index] = proposer.backend._gen_argmax_token_ids(draft_output) + + proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids[0][:real_verify_rows0] + proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids[1][:real_verify_rows1] + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py new file mode 100644 index 0000000000..642dc8848f --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -0,0 +1,82 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( + build_dp_chained_mtp_draft_state_from_prefill_overlap, + propose_next_dp_chained_mtp_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + build_chained_mtp_draft_state_from_prefill, + propose_next_chained_mtp, +) + + +class DpOverlapVanillaWithAttProposer(BaseDpOverlapProposer): + """DP ``vanilla_with_att`` proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_chained_mtp_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def build_draft_state_from_prefill_overlap( + self, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + ) -> None: + build_dp_chained_mtp_draft_state_from_prefill_overlap( + self, + target_model_input0, + target_model_output0, + next_token_ids0, + target_model_input1, + target_model_output1, + next_token_ids1, + ) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) + + def propose_next_overlap( + self, + main_model_input0: ModelInput, + main_model_output0: ModelOutput, + next_token_ids0: torch.Tensor, + real_verify_rows0: int, + accept_len0: torch.Tensor | None, + main_model_input1: ModelInput, + main_model_output1: ModelOutput, + next_token_ids1: torch.Tensor, + real_verify_rows1: int, + accept_len1: torch.Tensor | None, + draft_step: int, + ) -> SpecProposal: + return propose_next_dp_chained_mtp_overlap( + self, + main_model_input0, + main_model_output0, + next_token_ids0, + real_verify_rows0, + main_model_input1, + main_model_output1, + next_token_ids1, + real_verify_rows1, + draft_step, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py index 0a00c771a7..e29812e24d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py @@ -1 +1,9 @@ -"""DP-specific MTP planner implementations.""" +from lightllm.server.router.model_infer.mtp_speculative.dp_planner.base import BaseDpPlanner +from lightllm.server.router.model_infer.mtp_speculative.dp_planner.fixed import FixedDpPlanner + + +def build_dp_planner(*, backend) -> BaseDpPlanner: + return FixedDpPlanner(draft_step=backend.max_draft_step) + + +__all__ = ["BaseDpPlanner", "FixedDpPlanner", "build_dp_planner"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/base.py new file mode 100644 index 0000000000..bd20e39310 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/base.py @@ -0,0 +1,11 @@ +from abc import ABC, abstractmethod + + +class BaseDpPlanner(ABC): + """普通 DP draft 配置的基础规划接口。""" + + @abstractmethod + def get_draft_step(self) -> int: + """返回当前 DP decode proposal 使用的 draft step。""" + + raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/fixed.py new file mode 100644 index 0000000000..3fdf4a7b63 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/fixed.py @@ -0,0 +1,11 @@ +from lightllm.server.router.model_infer.mtp_speculative.dp_planner.base import BaseDpPlanner + + +class FixedDpPlanner(BaseDpPlanner): + """为 DP 固定 shape decode 返回启动时配置的 draft step。""" + + def __init__(self, draft_step: int) -> None: + self.draft_step = int(draft_step) + + def get_draft_step(self) -> int: + return self.draft_step diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py index 8274fbb5a4..4053a4ee27 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py @@ -1 +1,38 @@ -"""DP-specific MTP proposer implementations.""" +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer + + +def build_dp_spec_proposer(*, spec_mode: str, backend, enable_dynmaic_mtp: bool) -> "BaseDpProposer": + if spec_mode == "vanilla_with_att": + from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_with_att import ( + DpVanillaWithAttProposer, + ) + + return DpVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "vanilla_no_att": + from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_no_att import ( + DpVanillaNoAttProposer, + ) + + return DpVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "eagle_with_att": + from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_with_att import ( + DpEagleWithAttProposer, + ) + + return DpEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "eagle_no_att": + from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_no_att import DpEagleNoAttProposer + + return DpEagleNoAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "eagle3": + from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle3 import DpEagle3Proposer + + return DpEagle3Proposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + + raise ValueError(f"unsupported DP speculative mode: {spec_mode}") + + +__all__ = ["build_dp_spec_proposer"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/base.py new file mode 100644 index 0000000000..aada294da8 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/base.py @@ -0,0 +1,7 @@ +from abc import ABC + +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer + + +class BaseDpProposer(BaseSpecProposer, ABC): + """普通 DP prefill/decode proposer 的基础接口。""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py new file mode 100644 index 0000000000..6852354716 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py @@ -0,0 +1,44 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + propose_next_eagle, +) + + +class DpEagle3Proposer(BaseDpProposer): + """普通 DP ``eagle3`` proposer。""" + + def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: + return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + self._map_draft_token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py new file mode 100644 index 0000000000..d809363cb4 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py @@ -0,0 +1,41 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + propose_next_eagle, +) + + +class DpEagleNoAttProposer(BaseDpProposer): + """普通 DP ``eagle_no_att`` proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + lambda token_ids: token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py new file mode 100644 index 0000000000..b3cbbc9f12 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py @@ -0,0 +1,41 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + propose_next_eagle, +) + + +class DpEagleWithAttProposer(BaseDpProposer): + """普通 DP ``eagle_with_att`` proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + lambda token_ids: token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py new file mode 100644 index 0000000000..83f42a07c2 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -0,0 +1,32 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + build_chained_mtp_draft_state_from_prefill, + propose_next_chained_mtp, +) + + +class DpVanillaNoAttProposer(BaseDpProposer): + """普通 DP ``vanilla_no_att`` proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_chained_mtp_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py new file mode 100644 index 0000000000..c08e09b2fd --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -0,0 +1,32 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + build_chained_mtp_draft_state_from_prefill, + propose_next_chained_mtp, +) + + +class DpVanillaWithAttProposer(BaseDpProposer): + """普通 DP ``vanilla_with_att`` proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_chained_mtp_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 41dc7236a0..323a23e167 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -20,7 +20,7 @@ SpecDecodePlan, ) from lightllm.server.router.model_infer.mtp_speculative.proposers import build_spec_proposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal class SpecEngine: @@ -35,7 +35,7 @@ def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: self.backend = backend self.spec_mode = spec_mode self.enable_dynmaic_mtp = enable_dynmaic_mtp - self.proposer = build_spec_proposer( + self.proposer: BaseSpecProposer = build_spec_proposer( spec_mode=spec_mode, backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp, @@ -56,24 +56,6 @@ def build_draft_state_from_prefill( next_token_ids=next_token_ids, ) - def build_draft_state_from_prefill_overlap( - self, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, - ) -> None: - self.proposer.build_draft_state_from_prefill_overlap( - target_model_input0=target_model_input0, - target_model_output0=target_model_output0, - next_token_ids0=next_token_ids0, - target_model_input1=target_model_input1, - target_model_output1=target_model_output1, - next_token_ids1=next_token_ids1, - ) - # Decode planning and target verification. def plan_decode(self, model_input: ModelInput, decode_reqs: List) -> SpecDecodePlan: @@ -177,37 +159,6 @@ def propose_next( accept_len=accept_len, ) - def propose_next_overlap( - self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, - real_verify_rows1: int, - accept_len1: torch.Tensor, - draft_step: int, - ) -> SpecProposal: - # TODO: DP overlap currently disables dynamic drafting and always uses - # max_draft_step. Share the score-feedback path with propose_next when - # dynamic overlap is supported. - return self.proposer.propose_next_overlap( - main_model_input0=main_model_input0, - main_model_output0=main_model_output0, - next_token_ids0=next_token_ids0, - real_verify_rows0=real_verify_rows0, - accept_len0=accept_len0, - main_model_input1=main_model_input1, - main_model_output1=main_model_output1, - next_token_ids1=next_token_ids1, - real_verify_rows1=real_verify_rows1, - accept_len1=accept_len1, - draft_step=draft_step, - ) - def scatter_next_tokens( self, b_req_mtp_start_loc: torch.Tensor, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py index e4db7a9732..e633095e82 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py @@ -17,14 +17,24 @@ def build_spec_proposer(*, spec_mode: str, backend, enable_dynmaic_mtp: bool) -> from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer return Eagle3Proposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) - if spec_mode in ("eagle_with_att", "eagle_no_att"): - from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import EagleMTPProposer + if spec_mode == "eagle_with_att": + from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer - return EagleMTPProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) - if spec_mode in ("vanilla_with_att", "vanilla_no_att"): - from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_mtp import VanillaMTPProposer + return EagleWithAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "eagle_no_att": + from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import EagleNoAttProposer - return VanillaMTPProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + return EagleNoAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "vanilla_with_att": + from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import ( + VanillaWithAttProposer, + ) + + return VanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "vanilla_no_att": + from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer + + return VanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) raise ValueError(f"unsupported speculative mode: {spec_mode}") diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 1d275ea375..0f07062e49 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -1,5 +1,6 @@ from __future__ import annotations +from abc import ABC, abstractmethod from dataclasses import dataclass from typing import TYPE_CHECKING, Optional @@ -32,7 +33,7 @@ class SpecProposal: schedule_scores_cpu: Optional[torch.Tensor] = None -class BaseSpecProposer: +class BaseSpecProposer(ABC): """Base class for algorithm-specific draft proposal generation. A proposer owns the draft-side state transition. The target model gives it @@ -58,6 +59,7 @@ def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) return g_infer_context.req_manager.mem_manager.alloc(token_count) + @abstractmethod def build_draft_state_from_prefill( self, target_model_input: ModelInput, @@ -79,19 +81,7 @@ def build_draft_state_from_prefill( raise NotImplementedError - def build_draft_state_from_prefill_overlap( - self, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, - ) -> None: - """Build draft state from two overlapped target-prefill microbatches.""" - - raise NotImplementedError - + @abstractmethod def propose_next( self, main_model_input: ModelInput, @@ -111,21 +101,3 @@ def propose_next( """ raise NotImplementedError - - def propose_next_overlap( - self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, - real_verify_rows1: int, - accept_len1: torch.Tensor, - draft_step: int, - ) -> SpecProposal: - """Generate a proposal for two DP-overlapped microbatches.""" - - raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index 40f78681f2..1311bcaa26 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -3,17 +3,34 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block import ParallelBlockProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( + build_parallel_block_draft_input, + build_parallel_block_draft_state_from_prefill, + extend_parallel_block_draft_kv_cache, +) -class DFlashProposer(ParallelBlockProposer): +class DFlashProposer(BaseSpecProposer): """DFlash block-diffusion proposer. The drafter predicts a complete token block in one parallel forward from the accepted-tail anchor and mask-token positions. """ + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_parallel_block_draft_state_from_prefill( + proposer=self, + target_model_input=target_model_input, + target_model_output=target_model_output, + next_token_ids=next_token_ids, + ) + @torch.no_grad() def propose_next( self, @@ -36,13 +53,15 @@ def propose_next( # One accepted-tail anchor expands to a complete block-diffusion draft. accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_input, extra_mem_indexes_cpu = self.build_block_draft_input( + draft_input, extra_mem_indexes_cpu = build_parallel_block_draft_input( + proposer=self, main_model_input=main_model_input, next_token_ids=next_token_ids, accepted_tail_rows=accepted_tail_rows, request_count=request_count, ) - self.extend_draft_kv_cache( + extend_parallel_block_draft_kv_cache( + proposer=self, main_model_input=main_model_input, target_hidden=main_model_output.mtp_collector.spec_hidden, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 2abd4ffd02..71f59c6ddb 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -4,11 +4,15 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal -from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block import ParallelBlockProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( + build_parallel_block_draft_input, + build_parallel_block_draft_state_from_prefill, + extend_parallel_block_draft_kv_cache, +) -class DSparkProposer(ParallelBlockProposer): +class DSparkProposer(BaseSpecProposer): """DSpark semi-autoregressive parallel-block proposer. A parallel DFlash-style backbone generates the block features in one pass; @@ -16,6 +20,19 @@ class DSparkProposer(ParallelBlockProposer): The confidence head supplies per-position scheduling scores. """ + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_parallel_block_draft_state_from_prefill( + proposer=self, + target_model_input=target_model_input, + target_model_output=target_model_output, + next_token_ids=next_token_ids, + ) + @torch.no_grad() def propose_next( self, @@ -46,13 +63,15 @@ def propose_next( ) accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_input, extra_mem_indexes_cpu = self.build_block_draft_input( + draft_input, extra_mem_indexes_cpu = build_parallel_block_draft_input( + proposer=self, main_model_input=main_model_input, next_token_ids=next_token_ids, accepted_tail_rows=accepted_tail_rows, request_count=request_count, ) - self.extend_draft_kv_cache( + extend_parallel_block_draft_kv_cache( + proposer=self, main_model_input=main_model_input, target_hidden=main_model_output.mtp_collector.spec_hidden, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py index 257a714118..8ffe8cd477 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py @@ -1,16 +1,51 @@ -from __future__ import annotations - import torch -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import AutoregressiveEagleProposer - +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + generate_eagle_token_ids, + generate_eagle_token_ids_and_prob, + propose_next_eagle, +) -class Eagle3Proposer(AutoregressiveEagleProposer): - """Eagle3 proposer. - It uses the shared autoregressive proposal flow with an additional - draft-to-target vocabulary mapping step. - """ +class Eagle3Proposer(BaseSpecProposer): + """带有 draft-to-target vocabulary 映射的 EAGLE3 proposer。""" def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) + + def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: + return generate_eagle_token_ids(self, model_output, self._map_draft_token_ids) + + def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): + return generate_eagle_token_ids_and_prob(self, model_output, self._map_draft_token_ids) + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + proposer=self, + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=draft_step, + accept_len=accept_len, + map_draft_token_ids=self._map_draft_token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py deleted file mode 100644 index 96c1b74699..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_mtp.py +++ /dev/null @@ -1,443 +0,0 @@ -from __future__ import annotations - -import copy - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal - - -class AutoregressiveEagleProposer(BaseSpecProposer): - """Shared autoregressive drafting flow for EAGLE-family proposers.""" - - def build_draft_state_from_prefill( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - next_token_ids: torch.Tensor, - ) -> None: - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - prepare_mtp_prefill_inputs( - model_input=target_model_input, - b_next_token_ids=next_token_ids, - mtp_draft_input_hiddens=target_model_output.mtp_collector.spec_hidden, - ) - self.backend.draft_models[0].forward(target_model_input) - - def build_draft_state_from_prefill_overlap( - self, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, - ) -> None: - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - prepare_mtp_prefill_inputs( - model_input=target_model_input0, - b_next_token_ids=next_token_ids0, - mtp_draft_input_hiddens=target_model_output0.mtp_collector.spec_hidden, - ) - prepare_mtp_prefill_inputs( - model_input=target_model_input1, - b_next_token_ids=next_token_ids1, - mtp_draft_input_hiddens=target_model_output1.mtp_collector.spec_hidden, - ) - self.backend.draft_models[0].microbatch_overlap_prefill(target_model_input0, target_model_input1) - - def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: - return draft_token_ids - - def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: - return self._map_draft_token_ids(self.backend._gen_argmax_token_ids(model_output)) - - def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): - draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(model_output) - return self._map_draft_token_ids(draft_token_ids), draft_token_probs - - def prepare_verify_extend_input( - self, - model_input: ModelInput, - input_ids: torch.Tensor, - target_hidden: torch.Tensor, - ) -> None: - model_input.is_prefill = True - model_input.total_token_num = model_input.batch_size - model_input.prefix_total_token_num = 0 - model_input.max_cache_len = max(0, int(model_input.max_cache_len or 0)) - model_input.input_ids = input_ids - model_input.mtp_draft_input_hiddens = target_hidden - model_input.b_ready_cache_len = model_input.b_seq_len - 1 - model_input.b_prefill_start_loc = torch.arange( - model_input.batch_size, - dtype=torch.int32, - device=input_ids.device, - ) - - @staticmethod - def _pad_step_mem_indexes( - real_mem_indexes: torch.Tensor, - request_capacity: int, - hold_mem_index: int, - ) -> torch.Tensor: - padded = torch.full( - (request_capacity,), - hold_mem_index, - dtype=real_mem_indexes.dtype, - device=real_mem_indexes.device, - ) - padded[: real_mem_indexes.shape[0]].copy_(real_mem_indexes) - return padded - - def propose_next( - self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> SpecProposal: - """Run the common autoregressive EAGLE proposal flow.""" - - verify_row_count = int(next_token_ids.shape[0]) - request_count = int(b_req_mtp_start_loc.shape[0]) - proposal_token_ids = next_token_ids.new_full( - (verify_row_count, draft_step + 1), - fill_value=1, - ) - proposal_token_ids[:, 0].copy_(next_token_ids) - collect_schedule_scores = self.enable_dynmaic_mtp - schedule_scores = ( - torch.zeros( - (verify_row_count, draft_step), - dtype=torch.float32, - device=next_token_ids.device, - ) - if collect_schedule_scores - else None - ) - if draft_step == 0: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - schedule_scores=schedule_scores, - ) - - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_model = self.backend.draft_models[0] - position_delta = main_model_input.b_position_delta - self.prepare_verify_extend_input( - model_input=main_model_input, - input_ids=next_token_ids, - target_hidden=main_model_output.mtp_collector.spec_hidden, - ) - extend_output = draft_model.forward(main_model_input) - - accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) - if collect_schedule_scores: - draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(accepted_tail_output) - schedule_scores[accepted_tail_rows, 0] = draft_token_probs.float() - else: - draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) - proposal_token_ids[accepted_tail_rows, 1] = draft_token_ids - draft_hidden = extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows) - - if draft_step == 1: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - schedule_scores=schedule_scores, - ) - - extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(request_count * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) - draft_seq_lens = main_model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1 - max_kv_seq_len = main_model_input.max_kv_seq_len - draft_input = copy.copy(main_model_input) - draft_input.is_prefill = False - draft_input.batch_size = request_count - draft_input.b_req_idx = main_model_input.b_req_idx.index_select(0, accepted_tail_rows) - draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) - draft_input.b_seq_len = draft_seq_lens - draft_input.b_position_delta = ( - position_delta.index_select(0, accepted_tail_rows) if position_delta is not None else None - ) - draft_input.b_mark_shared_group = torch.ones_like(draft_input.b_req_idx) - draft_input.b_shared_seq_len = None - if len(draft_input.multimodal_params) != request_count: - empty_multimodal_params = {"images": [], "audios": []} - draft_input.multimodal_params = [empty_multimodal_params] * request_count - - for step in range(1, draft_step): - mem_start = (step - 1) * request_count - draft_input.input_ids = draft_token_ids - draft_input.mtp_draft_input_hiddens = draft_hidden - draft_input.mem_indexes = extra_mem_indexes[mem_start : mem_start + request_count] - draft_input.max_kv_seq_len = max_kv_seq_len + step - draft_input.total_token_num = request_count * draft_input.max_kv_seq_len - draft_output = draft_model.forward(draft_input) - if collect_schedule_scores: - draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(draft_output) - schedule_scores[accepted_tail_rows, step] = draft_token_probs.float() - else: - draft_token_ids = self._gen_argmax_token_ids(draft_output) - proposal_token_ids[accepted_tail_rows, step + 1] = draft_token_ids - draft_hidden = draft_output.mtp_collector.spec_hidden - draft_seq_lens.add_(1) - - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, - schedule_scores=schedule_scores, - ) - - def propose_next_overlap( - self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, - real_verify_rows1: int, - accept_len1: torch.Tensor, - draft_step: int, - ) -> SpecProposal: - """Run autoregressive EAGLE drafting with DP microbatch overlap. - - Target verification remains physically padded to ``B * (K + 1)`` for - DP collectives. Draft state is corrected once over those verify rows; - autoregressive drafting then runs one row per logical request (plus - HOLD rows required to keep the two microbatches shape-compatible). - """ - - verify_width = self.backend.max_draft_step + 1 - model_inputs = (main_model_input0, main_model_input1) - model_outputs = (main_model_output0, main_model_output1) - next_token_ids_by_batch = (next_token_ids0, next_token_ids1) - real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) - accept_lens_by_batch = (accept_len0, accept_len1) - position_deltas_by_batch = tuple(model_input.b_position_delta for model_input in model_inputs) - max_kv_seq_lens_by_batch = tuple(model_input.max_kv_seq_len for model_input in model_inputs) - - request_capacities_by_batch = [] - real_request_counts = [] - accepted_tail_rows_by_batch = [] - for model_input, model_output, token_ids, real_verify_row_count, accept_len in zip( - model_inputs, - model_outputs, - next_token_ids_by_batch, - real_verify_row_counts, - accept_lens_by_batch, - ): - request_capacity = model_input.batch_size // verify_width - real_request_count = real_verify_row_count // verify_width - starts = torch.arange( - 0, - model_input.batch_size, - verify_width, - dtype=torch.int32, - device=token_ids.device, - ) - accepted_tail_rows = (starts + accept_len - 1).long() - request_capacities_by_batch.append(request_capacity) - real_request_counts.append(real_request_count) - accepted_tail_rows_by_batch.append(accepted_tail_rows) - self.prepare_verify_extend_input( - model_input=model_input, - input_ids=token_ids, - target_hidden=model_output.mtp_collector.spec_hidden, - ) - - verify_row_count = real_verify_rows0 + real_verify_rows1 - proposal_token_ids = next_token_ids0.new_full( - (verify_row_count, draft_step + 1), - fill_value=1, - ) - proposal_token_ids[:real_verify_rows0, 0].copy_(next_token_ids0[:real_verify_rows0]) - proposal_token_ids[real_verify_rows0:, 0].copy_(next_token_ids1[:real_verify_rows1]) - - draft_model = self.backend.draft_models[0] - extend_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) - - draft_token_ids_by_batch = [] - draft_hiddens_by_batch = [] - draft_seq_lens_by_batch = [] - draft_req_indices_by_batch = [] - proposal_rows_by_batch = [] - proposal_row_offsets = (0, real_verify_rows0) - for batch_index, (model_input, extend_output, accepted_tail_rows, real_request_count) in enumerate( - zip(model_inputs, extend_outputs, accepted_tail_rows_by_batch, real_request_counts) - ): - accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) - draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) - draft_token_ids_by_batch.append(draft_token_ids) - draft_hiddens_by_batch.append(extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows)) - draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) - draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) - proposal_rows = accepted_tail_rows[:real_request_count] + proposal_row_offsets[batch_index] - proposal_rows_by_batch.append(proposal_rows) - proposal_token_ids[proposal_rows, 1] = draft_token_ids[:real_request_count] - - if draft_step == 1: - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - ) - - for batch_index, model_input in enumerate(model_inputs): - model_input.is_prefill = False - model_input.batch_size = request_capacities_by_batch[batch_index] - model_input.b_req_idx = draft_req_indices_by_batch[batch_index] - model_input.b_mtp_index = torch.zeros_like(model_input.b_req_idx) - model_input.b_seq_len = draft_seq_lens_by_batch[batch_index] - model_input.b_position_delta = ( - position_deltas_by_batch[batch_index].index_select(0, accepted_tail_rows_by_batch[batch_index]) - if position_deltas_by_batch[batch_index] is not None - else None - ) - model_input.b_mark_shared_group = torch.ones_like(model_input.b_req_idx) - model_input.b_shared_seq_len = None - if len(model_input.multimodal_params) != model_input.batch_size: - empty_multimodal_params = {"images": [], "audios": []} - model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size - - total_real_request_count = sum(real_request_counts) - extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(total_real_request_count * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) - hold_mem_index = self.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX - - for step in range(1, draft_step): - mem_start = (step - 1) * total_real_request_count - step_mem_indexes = extra_mem_indexes[mem_start : mem_start + total_real_request_count] - real_mem_start = 0 - for batch_index, model_input in enumerate(model_inputs): - real_request_count = real_request_counts[batch_index] - real_mem_indexes = step_mem_indexes[real_mem_start : real_mem_start + real_request_count] - real_mem_start += real_request_count - padded_mem_indexes = self._pad_step_mem_indexes( - real_mem_indexes=real_mem_indexes, - request_capacity=request_capacities_by_batch[batch_index], - hold_mem_index=hold_mem_index, - ) - model_input.input_ids = draft_token_ids_by_batch[batch_index] - model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - model_input.mem_indexes = padded_mem_indexes - model_input.max_kv_seq_len = max_kv_seq_lens_by_batch[batch_index] + step - model_input.total_token_num = model_input.batch_size * model_input.max_kv_seq_len - - draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) - for batch_index, draft_output in enumerate(draft_outputs): - draft_token_ids = self._gen_argmax_token_ids(draft_output) - draft_token_ids_by_batch[batch_index] = draft_token_ids - draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden - draft_seq_lens_by_batch[batch_index].add_(1) - real_request_count = real_request_counts[batch_index] - proposal_token_ids[proposal_rows_by_batch[batch_index], step + 1] = draft_token_ids[:real_request_count] - - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, - ) - - -class EagleMTPProposer(AutoregressiveEagleProposer): - """Autoregressive EAGLE proposer backed by an MTP draft model. - - The draft model keeps a cache and repeatedly feeds the previous proposal - hidden back into the same draft model. - """ - - def propose_next_overlap( - self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, - real_verify_rows1: int, - accept_len1: torch.Tensor, - draft_step: int, - ) -> SpecProposal: - """Run the fixed-layout Eagle MTP decode for two DP microbatches. - - DPEP overlap requires every rank to keep the target verify layout - throughout draft decode. Each logical request therefore remains a - contiguous group of ``draft_step + 1`` rows; accepted rows are not - compacted into a one-row autoregressive draft batch on this path. - """ - - verify_width = self.backend.max_draft_step + 1 - model_inputs = (main_model_input0, main_model_input1) - real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) - real_request_counts = tuple(row_count // verify_width for row_count in real_verify_row_counts) - request_capacities_by_batch = tuple(model_input.batch_size // verify_width for model_input in model_inputs) - total_real_request_count = sum(real_request_counts) - - proposal_token_ids = next_token_ids0.new_empty((sum(real_verify_row_counts), draft_step + 1)) - proposal_token_ids[:real_verify_rows0, 0] = next_token_ids0[:real_verify_rows0] - proposal_token_ids[real_verify_rows0:, 0] = next_token_ids1[:real_verify_rows1] - - extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(total_real_request_count * draft_step) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) - split = real_request_counts[0] * draft_step - extra_mem_indexes_by_batch = ( - extra_mem_indexes[:split], - extra_mem_indexes[split:], - ) - - draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] - draft_hiddens_by_batch = [ - main_model_output0.mtp_collector.spec_hidden, - main_model_output1.mtp_collector.spec_hidden, - ] - draft_model = self.backend.draft_models[0] - hold_mem_index = self.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX - - for step in range(draft_step): - for batch_index, model_input in enumerate(model_inputs): - model_input.input_ids = draft_token_ids_by_batch[batch_index] - model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - - draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) - - for batch_index, (model_input, draft_output) in enumerate(zip(model_inputs, draft_outputs)): - model_input.b_seq_len += 1 - model_input.max_kv_seq_len += 1 - - real_request_count = real_request_counts[batch_index] - mem_start = step * real_request_count - step_mem_indexes = extra_mem_indexes_by_batch[batch_index][mem_start : mem_start + real_request_count] - step_mem_indexes = self._pad_step_mem_indexes( - real_mem_indexes=step_mem_indexes, - request_capacity=request_capacities_by_batch[batch_index], - hold_mem_index=hold_mem_index, - ) - model_input.mem_indexes = torch.cat( - [ - model_input.mem_indexes.view(-1, verify_width)[:, 1:], - step_mem_indexes.view(-1, 1), - ], - dim=1, - ).view(-1) - - draft_token_ids_by_batch[batch_index] = self._gen_argmax_token_ids(draft_output) - draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden - - proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids_by_batch[0][:real_verify_rows0] - proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids_by_batch[1][:real_verify_rows1] - - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py new file mode 100644 index 0000000000..d74e8b87ee --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -0,0 +1,40 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + propose_next_eagle, +) + + +class EagleNoAttProposer(BaseSpecProposer): + """不使用 attention KV cache 的 EAGLE proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + proposer=self, + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=draft_step, + accept_len=accept_len, + map_draft_token_ids=lambda token_ids: token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py new file mode 100644 index 0000000000..66541e22e2 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -0,0 +1,186 @@ +"""EAGLE proposer 共享辅助函数。""" + +from __future__ import annotations + +import copy +from typing import Callable + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal + + +def build_eagle_draft_state_from_prefill( + proposer: BaseSpecProposer, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, +) -> None: + """使用 target prefill 输出初始化 EAGLE draft state。""" + + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + prepare_mtp_prefill_inputs( + model_input=target_model_input, + b_next_token_ids=next_token_ids, + mtp_draft_input_hiddens=target_model_output.mtp_collector.spec_hidden, + ) + proposer.backend.draft_models[0].forward(target_model_input) + + +def generate_eagle_token_ids( + proposer: BaseSpecProposer, + model_output: ModelOutput, + map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], +) -> torch.Tensor: + return map_draft_token_ids(proposer.backend._gen_argmax_token_ids(model_output)) + + +def generate_eagle_token_ids_and_prob( + proposer: BaseSpecProposer, + model_output: ModelOutput, + map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], +): + draft_token_ids, draft_token_probs = proposer.backend._gen_argmax_token_ids_and_prob(model_output) + return map_draft_token_ids(draft_token_ids), draft_token_probs + + +def prepare_eagle_verify_extend_input( + model_input: ModelInput, + input_ids: torch.Tensor, + target_hidden: torch.Tensor, +) -> None: + model_input.is_prefill = True + model_input.total_token_num = model_input.batch_size + model_input.prefix_total_token_num = 0 + model_input.max_cache_len = max(0, int(model_input.max_cache_len or 0)) + model_input.input_ids = input_ids + model_input.mtp_draft_input_hiddens = target_hidden + model_input.b_ready_cache_len = model_input.b_seq_len - 1 + model_input.b_prefill_start_loc = torch.arange( + model_input.batch_size, + dtype=torch.int32, + device=input_ids.device, + ) + + +def propose_next_eagle( + proposer: BaseSpecProposer, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None, + map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], +) -> SpecProposal: + """运行 EAGLE extend 后接单 token decode 的通用 proposal 流程。""" + + verify_row_count = int(next_token_ids.shape[0]) + request_count = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids = next_token_ids.new_full( + (verify_row_count, draft_step + 1), + fill_value=1, + ) + proposal_token_ids[:, 0].copy_(next_token_ids) + collect_schedule_scores = proposer.enable_dynmaic_mtp + schedule_scores = ( + torch.zeros( + (verify_row_count, draft_step), + dtype=torch.float32, + device=next_token_ids.device, + ) + if collect_schedule_scores + else None + ) + if draft_step == 0: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + schedule_scores=schedule_scores, + ) + + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + draft_model = proposer.backend.draft_models[0] + position_delta = main_model_input.b_position_delta + prepare_eagle_verify_extend_input( + model_input=main_model_input, + input_ids=next_token_ids, + target_hidden=main_model_output.mtp_collector.spec_hidden, + ) + extend_output = draft_model.forward(main_model_input) + + accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) + if collect_schedule_scores: + draft_token_ids, draft_token_probs = generate_eagle_token_ids_and_prob( + proposer=proposer, + model_output=accepted_tail_output, + map_draft_token_ids=map_draft_token_ids, + ) + schedule_scores[accepted_tail_rows, 0] = draft_token_probs.float() + else: + draft_token_ids = generate_eagle_token_ids( + proposer=proposer, + model_output=accepted_tail_output, + map_draft_token_ids=map_draft_token_ids, + ) + proposal_token_ids[accepted_tail_rows, 1] = draft_token_ids + draft_hidden = extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows) + + if draft_step == 1: + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + schedule_scores=schedule_scores, + ) + + extra_mem_indexes_cpu = proposer.alloc_extra_mem_indexes(request_count * (draft_step - 1)) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) + draft_seq_lens = main_model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1 + max_kv_seq_len = main_model_input.max_kv_seq_len + draft_input = copy.copy(main_model_input) + draft_input.is_prefill = False + draft_input.batch_size = request_count + draft_input.b_req_idx = main_model_input.b_req_idx.index_select(0, accepted_tail_rows) + draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) + draft_input.b_seq_len = draft_seq_lens + draft_input.b_position_delta = ( + position_delta.index_select(0, accepted_tail_rows) if position_delta is not None else None + ) + draft_input.b_mark_shared_group = torch.ones_like(draft_input.b_req_idx) + draft_input.b_shared_seq_len = None + if len(draft_input.multimodal_params) != request_count: + empty_multimodal_params = {"images": [], "audios": []} + draft_input.multimodal_params = [empty_multimodal_params] * request_count + + for step in range(1, draft_step): + mem_start = (step - 1) * request_count + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_input.mem_indexes = extra_mem_indexes[mem_start : mem_start + request_count] + draft_input.max_kv_seq_len = max_kv_seq_len + step + draft_input.total_token_num = request_count * draft_input.max_kv_seq_len + draft_output = draft_model.forward(draft_input) + if collect_schedule_scores: + draft_token_ids, draft_token_probs = generate_eagle_token_ids_and_prob( + proposer=proposer, + model_output=draft_output, + map_draft_token_ids=map_draft_token_ids, + ) + schedule_scores[accepted_tail_rows, step] = draft_token_probs.float() + else: + draft_token_ids = generate_eagle_token_ids( + proposer=proposer, + model_output=draft_output, + map_draft_token_ids=map_draft_token_ids, + ) + proposal_token_ids[accepted_tail_rows, step + 1] = draft_token_ids + draft_hidden = draft_output.mtp_collector.spec_hidden + draft_seq_lens.add_(1) + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, + schedule_scores=schedule_scores, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py new file mode 100644 index 0000000000..b66f809925 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -0,0 +1,40 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + build_eagle_draft_state_from_prefill, + propose_next_eagle, +) + + +class EagleWithAttProposer(BaseSpecProposer): + """使用 attention KV cache 的 EAGLE proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_eagle( + proposer=self, + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + draft_step=draft_step, + accept_len=accept_len, + map_draft_token_ids=lambda token_ids: token_ids, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py deleted file mode 100644 index 859a8f18a2..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block.py +++ /dev/null @@ -1,100 +0,0 @@ -from __future__ import annotations - -import copy - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer - - -class ParallelBlockProposer(BaseSpecProposer): - """Shared state and input preparation for parallel block drafters. - - DFlash and DSpark consume verified target output, extend the drafter KV - cache with target hidden rows, then generate a new block in one parallel - backbone forward. - """ - - @torch.no_grad() - def build_draft_state_from_prefill( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - next_token_ids: torch.Tensor, - ) -> None: - target_hidden = target_model_output.mtp_collector.spec_hidden - if target_hidden.numel() == 0: - return - - # Parallel block drafters consume target hidden states directly. - target_model_input.mtp_draft_input_hiddens = target_hidden - self.backend.draft_models[0].forward(target_model_input) - - def extend_draft_kv_cache(self, main_model_input: ModelInput, target_hidden: torch.Tensor) -> None: - # Target decode has finished; reuse its row-aligned input for draft KV commit. - main_model_input.total_token_num = main_model_input.batch_size - main_model_input.prefix_total_token_num = 0 - main_model_input.is_prefill = True - main_model_input.b_ready_cache_len = main_model_input.b_seq_len - 1 - main_model_input.b_prefill_start_loc = torch.arange( - main_model_input.batch_size, - dtype=torch.int32, - device=target_hidden.device, - ) - main_model_input.mtp_draft_input_hiddens = target_hidden - self.backend.draft_models[0].forward(main_model_input) - - def build_block_draft_input( - self, - main_model_input: ModelInput, - next_token_ids: torch.Tensor, - accepted_tail_rows: torch.Tensor, - request_count: int, - ): - draft_model = self.backend.draft_models[0] - block_size = int(draft_model.block_size) - extra_mem_indexes_cpu = self.alloc_extra_mem_indexes(request_count * block_size) - - block_input_ids = next_token_ids.new_full( - (request_count * block_size,), - fill_value=draft_model.mask_token_id, - ) - # Block input layout: [accepted token, mask, ..., mask]. The block - # logits become [base token + draft tokens] for target verification. - block_input_ids[::block_size] = next_token_ids.index_select(0, accepted_tail_rows) - - block_offsets = torch.arange( - block_size, - dtype=main_model_input.b_seq_len.dtype, - device=next_token_ids.device, - ) - draft_input = copy.copy(main_model_input) - draft_input.input_ids = block_input_ids - draft_input.total_token_num = draft_input.input_ids.shape[0] - draft_input.batch_size = draft_input.total_token_num - draft_input.max_q_seq_len = 1 - draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size - draft_input.b_req_idx = ( - main_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() - ) - draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) - # copy_kv_index_to_req and FA3 use these lengths to place scratch KV and - # compute the block cache length. - draft_input.b_seq_len = ( - (main_model_input.b_seq_len.index_select(0, accepted_tail_rows)[:, None] + block_offsets[None, :] + 1) - .reshape(-1) - .contiguous() - ) - # Position delta is request-level metadata shared by every block row. - draft_input.b_position_delta = ( - main_model_input.b_position_delta.index_select(0, accepted_tail_rows) - .repeat_interleave(block_size) - .contiguous() - ) - draft_input.mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) - draft_input.b_mark_shared_group = torch.zeros_like(draft_input.b_req_idx) - draft_input.b_mark_shared_group[block_size - 1 :: block_size] = block_size - empty_multimodal_params = {"images": [], "audios": []} - draft_input.multimodal_params = [empty_multimodal_params] * draft_input.batch_size - return draft_input, extra_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py new file mode 100644 index 0000000000..de6589a792 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py @@ -0,0 +1,97 @@ +"""Parallel-block proposer 共享辅助函数。""" + +from __future__ import annotations + +import copy + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer + + +@torch.no_grad() +def build_parallel_block_draft_state_from_prefill( + proposer: BaseSpecProposer, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, +) -> None: + """使用 target hidden 初始化 parallel-block drafter 的 KV state。""" + + target_hidden = target_model_output.mtp_collector.spec_hidden + if target_hidden.numel() == 0: + return + + target_model_input.mtp_draft_input_hiddens = target_hidden + proposer.backend.draft_models[0].forward(target_model_input) + + +def extend_parallel_block_draft_kv_cache( + proposer: BaseSpecProposer, + main_model_input: ModelInput, + target_hidden: torch.Tensor, +) -> None: + """提交本轮 target verify hidden,扩展 parallel-block drafter KV。""" + + main_model_input.total_token_num = main_model_input.batch_size + main_model_input.prefix_total_token_num = 0 + main_model_input.is_prefill = True + main_model_input.b_ready_cache_len = main_model_input.b_seq_len - 1 + main_model_input.b_prefill_start_loc = torch.arange( + main_model_input.batch_size, + dtype=torch.int32, + device=target_hidden.device, + ) + main_model_input.mtp_draft_input_hiddens = target_hidden + proposer.backend.draft_models[0].forward(main_model_input) + + +def build_parallel_block_draft_input( + proposer: BaseSpecProposer, + main_model_input: ModelInput, + next_token_ids: torch.Tensor, + accepted_tail_rows: torch.Tensor, + request_count: int, +): + """构建 parallel-block drafter 的 accepted-token + mask-token 输入。""" + + draft_model = proposer.backend.draft_models[0] + block_size = int(draft_model.block_size) + extra_mem_indexes_cpu = proposer.alloc_extra_mem_indexes(request_count * block_size) + + block_input_ids = next_token_ids.new_full( + (request_count * block_size,), + fill_value=draft_model.mask_token_id, + ) + block_input_ids[::block_size] = next_token_ids.index_select(0, accepted_tail_rows) + + block_offsets = torch.arange( + block_size, + dtype=main_model_input.b_seq_len.dtype, + device=next_token_ids.device, + ) + draft_input = copy.copy(main_model_input) + draft_input.input_ids = block_input_ids + draft_input.total_token_num = draft_input.input_ids.shape[0] + draft_input.batch_size = draft_input.total_token_num + draft_input.max_q_seq_len = 1 + draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size + draft_input.b_req_idx = ( + main_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() + ) + draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) + draft_input.b_seq_len = ( + (main_model_input.b_seq_len.index_select(0, accepted_tail_rows)[:, None] + block_offsets[None, :] + 1) + .reshape(-1) + .contiguous() + ) + draft_input.b_position_delta = ( + main_model_input.b_position_delta.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() + ) + draft_input.mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) + draft_input.b_mark_shared_group = torch.zeros_like(draft_input.b_req_idx) + draft_input.b_mark_shared_group[block_size - 1 :: block_size] = block_size + empty_multimodal_params = {"images": [], "audios": []} + draft_input.multimodal_params = [empty_multimodal_params] * draft_input.batch_size + return draft_input, extra_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py deleted file mode 100644 index a5a6f309b1..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_mtp.py +++ /dev/null @@ -1,113 +0,0 @@ -from __future__ import annotations - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal - - -class VanillaMTPProposer(BaseSpecProposer): - """Chained MTP proposer. - - Each draft depth uses an independent MTP module. Module i consumes the - hidden state produced by module i - 1 and predicts the next candidate. - """ - - def build_draft_state_from_prefill( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - next_token_ids: torch.Tensor, - ) -> None: - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_token_ids = next_token_ids - for draft_model in self.backend.draft_models: - prepare_mtp_prefill_inputs( - model_input=target_model_input, - b_next_token_ids=draft_token_ids, - mtp_draft_input_hiddens=draft_hidden, - ) - draft_output = draft_model.forward(target_model_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - - def build_draft_state_from_prefill_overlap( - self, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, - ) -> None: - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - draft_hiddens_by_batch = [ - target_model_output0.mtp_collector.spec_hidden, - target_model_output1.mtp_collector.spec_hidden, - ] - draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] - - for draft_model in self.backend.draft_models: - prepare_mtp_prefill_inputs( - model_input=target_model_input0, - b_next_token_ids=draft_token_ids_by_batch[0], - mtp_draft_input_hiddens=draft_hiddens_by_batch[0], - ) - prepare_mtp_prefill_inputs( - model_input=target_model_input1, - b_next_token_ids=draft_token_ids_by_batch[1], - mtp_draft_input_hiddens=draft_hiddens_by_batch[1], - ) - draft_outputs = draft_model.microbatch_overlap_prefill( - target_model_input0, - target_model_input1, - ) - for batch_index, draft_output in enumerate(draft_outputs): - draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden - draft_token_ids_by_batch[batch_index] = self.backend._gen_argmax_token_ids(draft_output) - - def propose_next( - self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> SpecProposal: - verify_row_count = int(next_token_ids.shape[0]) - draft_token_ids = next_token_ids - draft_hidden = main_model_output.mtp_collector.spec_hidden - proposal_token_ids = next_token_ids.new_empty((verify_row_count, draft_step + 1)) - proposal_token_ids[:, 0] = next_token_ids - schedule_scores = ( - torch.empty( - (verify_row_count, draft_step), - dtype=torch.float32, - device=next_token_ids.device, - ) - if self.enable_dynmaic_mtp - else None - ) - - for step in range(draft_step): - draft_model = self.backend.draft_models[step] - main_model_input.input_ids = draft_token_ids - main_model_input.mtp_draft_input_hiddens = draft_hidden - draft_output = draft_model.forward(main_model_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - if self.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) - schedule_scores[:, step] = draft_token_probs - else: - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - proposal_token_ids[:, step + 1] = draft_token_ids - - return SpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, - schedule_scores=schedule_scores, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py new file mode 100644 index 0000000000..12f5f48c61 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -0,0 +1,42 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + build_chained_mtp_draft_state_from_prefill, + propose_next_chained_mtp, +) + + +class VanillaNoAttProposer(BaseSpecProposer): + """不使用 attention KV cache 的 Vanilla chained MTP proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_chained_mtp_draft_state_from_prefill( + proposer=self, + target_model_input=target_model_input, + target_model_output=target_model_output, + next_token_ids=next_token_ids, + ) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_chained_mtp( + proposer=self, + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + draft_step=draft_step, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py new file mode 100644 index 0000000000..23718029f9 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py @@ -0,0 +1,75 @@ +"""Vanilla chained MTP proposer 共享辅助函数。""" + +from __future__ import annotations + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal + + +def build_chained_mtp_draft_state_from_prefill( + proposer: BaseSpecProposer, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, +) -> None: + """构建 Vanilla chained MTP 各级 draft model 的 prefill state。""" + + from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_token_ids = next_token_ids + for draft_model in proposer.backend.draft_models: + prepare_mtp_prefill_inputs( + model_input=target_model_input, + b_next_token_ids=draft_token_ids, + mtp_draft_input_hiddens=draft_hidden, + ) + draft_output = draft_model.forward(target_model_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + draft_token_ids = proposer.backend._gen_argmax_token_ids(draft_output) + + +def propose_next_chained_mtp( + proposer: BaseSpecProposer, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + draft_step: int, +) -> SpecProposal: + """依次运行 Vanilla chained MTP 模块并生成 proposal。""" + + verify_row_count = int(next_token_ids.shape[0]) + draft_token_ids = next_token_ids + draft_hidden = main_model_output.mtp_collector.spec_hidden + proposal_token_ids = next_token_ids.new_empty((verify_row_count, draft_step + 1)) + proposal_token_ids[:, 0] = next_token_ids + schedule_scores = ( + torch.empty( + (verify_row_count, draft_step), + dtype=torch.float32, + device=next_token_ids.device, + ) + if proposer.enable_dynmaic_mtp + else None + ) + + for step in range(draft_step): + draft_model = proposer.backend.draft_models[step] + main_model_input.input_ids = draft_token_ids + main_model_input.mtp_draft_input_hiddens = draft_hidden + draft_output = draft_model.forward(main_model_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + if proposer.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = proposer.backend._gen_argmax_token_ids_and_prob(draft_output) + schedule_scores[:, step] = draft_token_probs + else: + draft_token_ids = proposer.backend._gen_argmax_token_ids(draft_output) + proposal_token_ids[:, step + 1] = draft_token_ids + + return SpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=None, + schedule_scores=schedule_scores, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py new file mode 100644 index 0000000000..5807c5dfdf --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -0,0 +1,42 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + build_chained_mtp_draft_state_from_prefill, + propose_next_chained_mtp, +) + + +class VanillaWithAttProposer(BaseSpecProposer): + """使用 attention KV cache 的 Vanilla chained MTP proposer。""" + + def build_draft_state_from_prefill( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + next_token_ids: torch.Tensor, + ) -> None: + build_chained_mtp_draft_state_from_prefill( + proposer=self, + target_model_input=target_model_input, + target_model_output=target_model_output, + next_token_ids=next_token_ids, + ) + + def propose_next( + self, + main_model_input: ModelInput, + main_model_output: ModelOutput, + next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + draft_step: int, + accept_len: torch.Tensor | None = None, + ) -> SpecProposal: + return propose_next_chained_mtp( + proposer=self, + main_model_input=main_model_input, + main_model_output=main_model_output, + next_token_ids=next_token_ids, + draft_step=draft_step, + ) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py new file mode 100644 index 0000000000..14e2e82620 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -0,0 +1,234 @@ +from types import SimpleNamespace + +import torch + +from lightllm.server.router.model_infer.mode_backend.dp_backend.impl import DPChunkedPrefillBackend +from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend +from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ChunkedPrefillBackend +from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + + +class _RecordingSpecEngine: + def __init__(self): + self.propose_args = None + self.propose_overlap_args = None + self.scatter_args = None + + def propose_next(self, **kwargs): + self.propose_args = kwargs + token_ids = kwargs["next_token_ids"].new_zeros((16, 8)) + return SpecProposal( + token_ids=token_ids, + extra_mem_indexes_cpu=torch.tensor([123], dtype=torch.int32), + ) + + def propose_next_overlap(self, **kwargs): + self.propose_overlap_args = kwargs + row_count = kwargs["real_verify_rows0"] + kwargs["real_verify_rows1"] + token_ids = kwargs["next_token_ids0"].new_zeros((row_count, 8)) + return SpecProposal( + token_ids=token_ids, + extra_mem_indexes_cpu=torch.tensor([456], dtype=torch.int32), + ) + + def scatter_next_tokens(self, **kwargs): + self.scatter_args = kwargs + + +def test_backends_initialize_their_own_spec_engine(): + args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False) + backend = ChunkedPrefillBackend.__new__(ChunkedPrefillBackend) + backend.args = args + backend.max_draft_step = 2 + backend.init_spec_engine() + dp_backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + dp_backend.args = args + dp_backend.max_draft_step = 2 + dp_backend.enable_prefill_microbatch_overlap = False + dp_backend.enable_decode_microbatch_overlap = True + dp_backend.init_spec_engine() + + assert "spec_engine_class" not in ModeBackend.__dict__ + assert type(backend.spec_engine) is SpecEngine + assert type(dp_backend.spec_engine) is DPSpecEngine + assert type(dp_backend.dp_overlap_spec_engine) is DPOverlapSpecEngine + assert dp_backend.prefill_draft_engine is dp_backend.spec_engine + assert dp_backend.decode_draft_engine is dp_backend.dp_overlap_spec_engine + assert not issubclass(DPSpecEngine, SpecEngine) + assert not issubclass(DPOverlapSpecEngine, SpecEngine) + + +def test_dp_prefill_and_decode_select_overlap_engine_independently(): + args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False) + + for prefill_overlap in (False, True): + for decode_overlap in (False, True): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.args = args + backend.max_draft_step = 2 + backend.enable_prefill_microbatch_overlap = prefill_overlap + backend.enable_decode_microbatch_overlap = decode_overlap + backend.init_spec_engine() + + expected_prefill_engine = backend.dp_overlap_spec_engine if prefill_overlap else backend.spec_engine + expected_decode_engine = backend.dp_overlap_spec_engine if decode_overlap else backend.spec_engine + assert backend.prefill_draft_engine is expected_prefill_engine + assert backend.decode_draft_engine is expected_decode_engine + + +def test_padded_token_ids_support_empty_dp_rank(): + padded_token_ids = DPChunkedPrefillBackend._build_padded_next_token_ids( + token_ids=None, + batch_size=4, + copy_len=0, + device=torch.device("cpu"), + ) + + assert torch.equal(padded_token_ids, torch.zeros(4, dtype=torch.int64)) + + +def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.max_draft_step = 7 + backend.spec_engine = _RecordingSpecEngine() + backend.decode_draft_engine = backend.spec_engine + model_input = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + model_output = SimpleNamespace(spec_hidden=torch.randn(16, 4)) + next_token_ids = torch.arange(8, dtype=torch.int64) + real_start_locs = torch.tensor([0], dtype=torch.int32) + real_accept_len = torch.tensor([2], dtype=torch.int32) + + extra_mem = backend._draft_decode_eagle( + model_input=model_input, + model_output=model_output, + next_token_ids=next_token_ids, + b_req_mtp_start_loc=real_start_locs, + mtp_accept_len=real_accept_len, + req_num=8, + ) + + propose_args = backend.spec_engine.propose_args + assert propose_args["next_token_ids"].shape == (16,) + assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) + assert torch.equal(propose_args["b_req_mtp_start_loc"], torch.tensor([0, 8], dtype=torch.int32)) + assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) + assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (8, 8) + assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) + + +def test_dp_vanilla_uses_dp_engine_proposer(): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.max_draft_step = 7 + backend.spec_engine = _RecordingSpecEngine() + backend.decode_draft_engine = backend.spec_engine + model_input = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + next_token_ids = torch.arange(8, dtype=torch.int64) + + extra_mem = backend._draft_decode_vanilla( + model_input=model_input, + model_output=SimpleNamespace(), + next_token_ids=next_token_ids, + b_req_mtp_start_loc=torch.arange(8, dtype=torch.int32), + mtp_accept_len=torch.ones(8, dtype=torch.int32), + req_num=8, + ) + + propose_args = backend.spec_engine.propose_args + assert propose_args["next_token_ids"].shape == (16,) + assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) + assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (8, 8) + assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) + + +def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.max_draft_step = 7 + backend.spec_engine = _RecordingSpecEngine() + backend.dp_overlap_spec_engine = _RecordingSpecEngine() + backend.decode_draft_engine = backend.dp_overlap_spec_engine + model_input0 = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + model_input1 = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + model_output0 = SimpleNamespace(spec_hidden=torch.randn(16, 4)) + model_output1 = SimpleNamespace(spec_hidden=torch.randn(16, 4)) + next_token_ids = torch.arange(24, dtype=torch.int64) + b_req_idx = torch.arange(24, dtype=torch.int32) + start_locs = torch.tensor([0, 8, 16], dtype=torch.int32) + accept_len = torch.tensor([2, 3, 4], dtype=torch.int32) + + extra_mem = backend._draft_decode_eagle_overlap( + model_input0=model_input0, + model_output0=model_output0, + model_input1=model_input1, + model_output1=model_output1, + b_req_idx=b_req_idx, + next_token_ids=next_token_ids, + mtp_accept_len=accept_len, + b_req_mtp_start_loc=start_locs, + req_num0=8, + req_num1=16, + ) + + propose_args = backend.dp_overlap_spec_engine.propose_overlap_args + assert propose_args["next_token_ids0"].shape == (16,) + assert propose_args["next_token_ids1"].shape == (16,) + assert propose_args["real_verify_rows0"] == 8 + assert propose_args["real_verify_rows1"] == 16 + assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) + assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) + assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (24, 8) + assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) + + +def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.max_draft_step = 7 + backend.spec_engine = _RecordingSpecEngine() + backend.dp_overlap_spec_engine = _RecordingSpecEngine() + backend.decode_draft_engine = backend.dp_overlap_spec_engine + model_input0 = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + model_input1 = SimpleNamespace( + batch_size=16, + b_req_idx=torch.arange(16, dtype=torch.int32), + ) + next_token_ids = torch.arange(24, dtype=torch.int64) + + extra_mem = backend._draft_decode_vanilla_overlap( + model_input0=model_input0, + model_output0=SimpleNamespace(), + model_input1=model_input1, + model_output1=SimpleNamespace(), + b_req_idx=torch.arange(24, dtype=torch.int32), + next_token_ids=next_token_ids, + mtp_accept_len=torch.tensor([2, 3, 4], dtype=torch.int32), + b_req_mtp_start_loc=torch.tensor([0, 8, 16], dtype=torch.int32), + req_num0=8, + req_num1=16, + ) + + propose_args = backend.dp_overlap_spec_engine.propose_overlap_args + assert propose_args["next_token_ids0"].shape == (16,) + assert propose_args["next_token_ids1"].shape == (16,) + assert propose_args["real_verify_rows0"] == 8 + assert propose_args["real_verify_rows1"] == 16 + assert propose_args["accept_len0"] is None + assert propose_args["accept_len1"] is None + assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (24, 8) + assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py deleted file mode 100644 index 1aef6b74b6..0000000000 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_spec_engine.py +++ /dev/null @@ -1,118 +0,0 @@ -from types import SimpleNamespace - -import torch - -from lightllm.server.router.model_infer.mode_backend.dp_backend.impl import DPChunkedPrefillBackend -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal - - -class _RecordingSpecEngine: - def __init__(self): - self.propose_args = None - self.propose_overlap_args = None - self.scatter_args = None - - def propose_next(self, **kwargs): - self.propose_args = kwargs - token_ids = kwargs["next_token_ids"].new_zeros((16, 8)) - return SpecProposal( - token_ids=token_ids, - extra_mem_indexes_cpu=torch.tensor([123], dtype=torch.int32), - ) - - def propose_next_overlap(self, **kwargs): - self.propose_overlap_args = kwargs - row_count = kwargs["real_verify_rows0"] + kwargs["real_verify_rows1"] - token_ids = kwargs["next_token_ids0"].new_zeros((row_count, 8)) - return SpecProposal( - token_ids=token_ids, - extra_mem_indexes_cpu=torch.tensor([456], dtype=torch.int32), - ) - - def scatter_next_tokens(self, **kwargs): - self.scatter_args = kwargs - - -def test_padded_token_ids_support_empty_dp_rank(): - padded_token_ids = DPChunkedPrefillBackend._build_padded_next_token_ids( - token_ids=None, - batch_size=4, - copy_len=0, - device=torch.device("cpu"), - ) - - assert torch.equal(padded_token_ids, torch.zeros(4, dtype=torch.int64)) - - -def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(): - backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) - backend.max_draft_step = 7 - backend.spec_engine = _RecordingSpecEngine() - model_input = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), - ) - model_output = SimpleNamespace(spec_hidden=torch.randn(16, 4)) - next_token_ids = torch.arange(8, dtype=torch.int64) - real_start_locs = torch.tensor([0], dtype=torch.int32) - real_accept_len = torch.tensor([2], dtype=torch.int32) - - extra_mem = backend._draft_decode_eagle( - model_input=model_input, - model_output=model_output, - next_token_ids=next_token_ids, - b_req_mtp_start_loc=real_start_locs, - mtp_accept_len=real_accept_len, - req_num=8, - ) - - propose_args = backend.spec_engine.propose_args - assert propose_args["next_token_ids"].shape == (16,) - assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) - assert torch.equal(propose_args["b_req_mtp_start_loc"], torch.tensor([0, 8], dtype=torch.int32)) - assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) - assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (8, 8) - assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) - - -def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(): - backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) - backend.max_draft_step = 7 - backend.spec_engine = _RecordingSpecEngine() - model_input0 = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), - ) - model_input1 = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), - ) - model_output0 = SimpleNamespace(spec_hidden=torch.randn(16, 4)) - model_output1 = SimpleNamespace(spec_hidden=torch.randn(16, 4)) - next_token_ids = torch.arange(24, dtype=torch.int64) - b_req_idx = torch.arange(24, dtype=torch.int32) - start_locs = torch.tensor([0, 8, 16], dtype=torch.int32) - accept_len = torch.tensor([2, 3, 4], dtype=torch.int32) - - extra_mem = backend._draft_decode_eagle_overlap( - model_input0=model_input0, - model_output0=model_output0, - model_input1=model_input1, - model_output1=model_output1, - b_req_idx=b_req_idx, - next_token_ids=next_token_ids, - mtp_accept_len=accept_len, - b_req_mtp_start_loc=start_locs, - req_num0=8, - req_num1=16, - ) - - propose_args = backend.spec_engine.propose_overlap_args - assert propose_args["next_token_ids0"].shape == (16,) - assert propose_args["next_token_ids1"].shape == (16,) - assert propose_args["real_verify_rows0"] == 8 - assert propose_args["real_verify_rows1"] == 16 - assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) - assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) - assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (24, 8) - assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index a0c663a08c..0aa2c0378f 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -3,11 +3,11 @@ import torch from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import ( - AutoregressiveEagleProposer, - EagleMTPProposer, +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import DpOverlapEagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_with_att import ( + DpOverlapEagleWithAttProposer, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer class _DraftModel: @@ -39,6 +39,9 @@ def microbatch_overlap_decode(self, input0, input1): for model_input in (input0, input1) ) + def map_draft_vocab_to_main_vocab(self, token_ids): + return token_ids + def _target_input(batch_size): return SimpleNamespace( @@ -71,7 +74,7 @@ def test_overlap_eagle_keeps_fixed_verify_layout(): ), _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) - proposer = EagleMTPProposer(backend=backend, enable_dynmaic_mtp=False) + proposer = DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=False) proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) model_input0 = _target_input(batch_size=6) model_input1 = _target_input(batch_size=6) @@ -118,7 +121,7 @@ def test_autoregressive_eagle_reuses_overlap_inputs(): ), _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) - proposer = AutoregressiveEagleProposer(backend=backend, enable_dynmaic_mtp=False) + proposer = DpOverlapEagle3Proposer(backend=backend, enable_dynmaic_mtp=False) proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) model_input0 = _target_input(batch_size=6) model_input1 = _target_input(batch_size=6) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 32d7c88ebe..23462f40f1 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -5,6 +5,42 @@ import torch from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine +from lightllm.server.router.model_infer.mtp_speculative.dp_planner import ( + BaseDpPlanner, + FixedDpPlanner, + build_dp_planner, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner import ( + BaseDpOverlapPlanner, + FixedDpOverlapPlanner, + build_dp_overlap_planner, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers import build_dp_spec_proposer +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle3 import DpEagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_no_att import DpEagleNoAttProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_with_att import DpEagleWithAttProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_no_att import DpVanillaNoAttProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_with_att import ( + DpVanillaWithAttProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers import ( + build_dp_overlap_spec_proposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import DpOverlapEagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_no_att import ( + DpOverlapEagleNoAttProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_with_att import ( + DpOverlapEagleWithAttProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_no_att import ( + DpOverlapVanillaNoAttProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_with_att import ( + DpOverlapVanillaWithAttProposer, +) from lightllm.server.router.model_infer.mtp_speculative.planner import ( BaseMtpPlanner, DSparkPlanner, @@ -13,16 +49,15 @@ SpecDecodePlan, ) from lightllm.server.router.model_infer.mtp_speculative.planner.base import _InferCostMsTable -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers import build_spec_proposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_mtp import ( - AutoregressiveEagleProposer, - EagleMTPProposer, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block import ParallelBlockProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_mtp import VanillaMTPProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import EagleNoAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import VanillaWithAttProposer def build_lightspec_planner( @@ -53,9 +88,9 @@ def build_dspark_planner(max_draft_step: int = 3, block_size: int = 3): def build_decode_reqs(req_num: int, req_num_with_proposals: int | None = None): if req_num_with_proposals is None: req_num_with_proposals = req_num - return [SimpleNamespace(cur_output_len=2)] * req_num_with_proposals + [ - SimpleNamespace(cur_output_len=1) - ] * (req_num - req_num_with_proposals) + return [SimpleNamespace(cur_output_len=2)] * req_num_with_proposals + [SimpleNamespace(cur_output_len=1)] * ( + req_num - req_num_with_proposals + ) def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): @@ -81,6 +116,22 @@ def test_fixed_planner_returns_static_plan(): assert not plan.skip_verify_sync +def test_dp_planner_returns_fixed_backend_draft_step(): + planner = build_dp_planner(backend=SimpleNamespace(max_draft_step=4)) + + assert type(planner) is FixedDpPlanner + assert isinstance(planner, BaseDpPlanner) + assert planner.get_draft_step() == 4 + + +def test_dp_overlap_planner_returns_fixed_backend_draft_step(): + planner = build_dp_overlap_planner(backend=SimpleNamespace(max_draft_step=5)) + + assert type(planner) is FixedDpOverlapPlanner + assert isinstance(planner, BaseDpOverlapPlanner) + assert planner.get_draft_step() == 5 + + def test_infer_cost_candidates_include_feasible_boundaries(): costs = _InferCostMsTable() costs.update(batch_size=4, infer_cost_ms=1.0) @@ -135,16 +186,90 @@ def test_dynamic_planner_registers_cuda_graph_costs_from_backend(): assert planner.draft_infer_costs.estimate(4) == 0.3 -def test_proposer_families_use_drafter_standard_abstractions(): - assert issubclass(EagleMTPProposer, AutoregressiveEagleProposer) - assert issubclass(Eagle3Proposer, AutoregressiveEagleProposer) - assert issubclass(DFlashProposer, ParallelBlockProposer) - assert issubclass(DSparkProposer, ParallelBlockProposer) - assert not issubclass(DSparkProposer, DFlashProposer) +def test_each_mode_proposer_only_inherits_its_base_interface(): + proposer_types = ( + VanillaWithAttProposer, + VanillaNoAttProposer, + EagleWithAttProposer, + EagleNoAttProposer, + Eagle3Proposer, + DFlashProposer, + DSparkProposer, + ) + dp_proposer_types = ( + DpVanillaWithAttProposer, + DpVanillaNoAttProposer, + DpEagleWithAttProposer, + DpEagleNoAttProposer, + DpEagle3Proposer, + ) + dp_overlap_proposer_types = ( + DpOverlapVanillaWithAttProposer, + DpOverlapVanillaNoAttProposer, + DpOverlapEagleWithAttProposer, + DpOverlapEagleNoAttProposer, + DpOverlapEagle3Proposer, + ) + + for proposer_type in proposer_types: + assert proposer_type.__bases__ == (BaseSpecProposer,) + for proposer_type in dp_proposer_types: + assert proposer_type.__bases__ == (BaseDpProposer,) + for proposer_type in dp_overlap_proposer_types: + assert proposer_type.__bases__ == (BaseDpOverlapProposer,) + + +def test_each_mtp_mode_builds_its_own_proposer(): + backend = SimpleNamespace() + proposer_types = { + "vanilla_with_att": VanillaWithAttProposer, + "vanilla_no_att": VanillaNoAttProposer, + "eagle_with_att": EagleWithAttProposer, + "eagle_no_att": EagleNoAttProposer, + "eagle3": Eagle3Proposer, + "dflash": DFlashProposer, + "dspark": DSparkProposer, + } + + for spec_mode, proposer_type in proposer_types.items(): + proposer = build_spec_proposer(spec_mode=spec_mode, backend=backend, enable_dynmaic_mtp=False) + assert type(proposer) is proposer_type + + +def test_each_supported_dp_mtp_mode_builds_its_own_dp_proposer(): + backend = SimpleNamespace() + proposer_types = { + "vanilla_with_att": DpVanillaWithAttProposer, + "vanilla_no_att": DpVanillaNoAttProposer, + "eagle_with_att": DpEagleWithAttProposer, + "eagle_no_att": DpEagleNoAttProposer, + "eagle3": DpEagle3Proposer, + } + + for spec_mode, proposer_type in proposer_types.items(): + proposer = build_dp_spec_proposer(spec_mode=spec_mode, backend=backend, enable_dynmaic_mtp=False) + assert type(proposer) is proposer_type + assert isinstance(proposer, BaseDpProposer) + + +def test_each_supported_dp_overlap_mtp_mode_builds_its_own_proposer(): + backend = SimpleNamespace() + proposer_types = { + "vanilla_with_att": DpOverlapVanillaWithAttProposer, + "vanilla_no_att": DpOverlapVanillaNoAttProposer, + "eagle_with_att": DpOverlapEagleWithAttProposer, + "eagle_no_att": DpOverlapEagleNoAttProposer, + "eagle3": DpOverlapEagle3Proposer, + } + + for spec_mode, proposer_type in proposer_types.items(): + proposer = build_dp_overlap_spec_proposer(spec_mode=spec_mode, backend=backend, enable_dynmaic_mtp=False) + assert type(proposer) is proposer_type + assert isinstance(proposer, BaseDpOverlapProposer) def test_eagle_proposer_skips_draft_forward_for_zero_steps(): - proposer = EagleMTPProposer( + proposer = EagleNoAttProposer( backend=SimpleNamespace(draft_models=[]), enable_dynmaic_mtp=True, ) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py new file mode 100644 index 0000000000..8e738a9b7f --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -0,0 +1,65 @@ +from types import SimpleNamespace + +import torch + +from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_with_att import ( + DpOverlapVanillaWithAttProposer, +) + + +class _DraftModel: + def __init__(self): + self.decode_batch_sizes = [] + + def microbatch_overlap_decode(self, input0, input1): + self.decode_batch_sizes.append((input0.batch_size, input1.batch_size)) + return tuple( + ModelOutput( + logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), + ) + for model_input in (input0, input1) + ) + + +def test_dp_vanilla_proposer_owns_overlap_decode(): + draft_models = [_DraftModel(), _DraftModel()] + backend = SimpleNamespace( + draft_models=draft_models, + _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), + ) + proposer = DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) + model_input0 = SimpleNamespace(batch_size=4) + model_input1 = SimpleNamespace(batch_size=4) + + proposal = proposer.propose_next_overlap( + main_model_input0=model_input0, + main_model_output0=ModelOutput( + logits=torch.empty((4, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((4, 2))), + ), + next_token_ids0=torch.tensor([10, 11, 0, 0], dtype=torch.int64), + real_verify_rows0=2, + accept_len0=None, + main_model_input1=model_input1, + main_model_output1=ModelOutput( + logits=torch.empty((4, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((4, 2))), + ), + next_token_ids1=torch.tensor([20, 21, 22, 0], dtype=torch.int64), + real_verify_rows1=3, + accept_len1=None, + draft_step=2, + ) + + assert proposal.token_ids.tolist() == [ + [10, 0, 0], + [11, 1, 1], + [20, 0, 0], + [21, 1, 1], + [22, 2, 2], + ] + assert proposal.extra_mem_indexes_cpu is None + assert draft_models[0].decode_batch_sizes == [(4, 4)] + assert draft_models[1].decode_batch_sizes == [(4, 4)] From addeab42e6ca6187738969b9a672ce562709d267 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 02:15:12 +0000 Subject: [PATCH 046/103] refactor: unify fixed and dynamic MTP decode plans --- .../model_infer/mtp_speculative/engine.py | 16 +++---- .../mtp_speculative/planner/base.py | 32 +++++++------- .../mtp_speculative/planner/dspark.py | 12 ++---- .../mtp_speculative/planner/fixed.py | 5 ++- .../mtp_speculative/planner/lightspec.py | 13 ++---- .../mtp_speculative/test_planner.py | 43 +++++++++++-------- 6 files changed, 57 insertions(+), 64 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 323a23e167..ebf080ccd7 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -61,9 +61,10 @@ def build_draft_state_from_prefill( def plan_decode(self, model_input: ModelInput, decode_reqs: List) -> SpecDecodePlan: """Return the fixed or dynamic speculative plan for one decode iteration.""" + assert decode_reqs, "non-DP speculative decode requires at least one request" return self.planner.plan( decode_reqs=decode_reqs, - original_batch_size=model_input.batch_size, + origin_batch_size=model_input.batch_size, ) def prepare_decode_model_input( @@ -72,12 +73,10 @@ def prepare_decode_model_input( req_num: int, plan: SpecDecodePlan, ) -> Tuple[ModelInput, Optional[AsyncPinnedCpuTensor]]: - """Apply target verify-row compaction when the dynamic planner selects it.""" + """Apply target verify-row compaction when the planned batch is smaller.""" - if not plan.is_dynamic: - return model_input, None - - if plan.dynamic_batch_size == model_input.batch_size: + assert model_input.batch_size == plan.origin_batch_size + if plan.dynamic_batch_size == plan.origin_batch_size: return model_input, None from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_mtp_model_input @@ -190,10 +189,7 @@ def update_planner_feedback( req_num: int, accept_lengths_cpu: torch.Tensor, ) -> None: - """Feed iteration-level observations into the dynamic planner.""" - - if not self.enable_dynmaic_mtp: - return + """Feed iteration-level observations into the current planner.""" self.planner.update_feedback( plan=plan, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py index 7fc6aaa33d..7802421632 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -2,7 +2,7 @@ from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import List, Optional +from typing import List from sortedcontainers import SortedDict @@ -11,18 +11,21 @@ class SpecDecodePlan: """Planner decision for one target decode iteration. - Fixed scheduling uses the full speculative-expanded target batch: - - dynamic_batch_size is None - - draft_step == max_draft_step + ``origin_batch_size`` records the physical target row count before planning, + while ``dynamic_batch_size`` is the row count selected for this target + forward. Fixed scheduling sets both values to the same size; dynamic + scheduling may reduce ``dynamic_batch_size`` before the forward. Therefore + both modes share the same plan representation without using ``None`` as a + mode marker. - Dynamic speculative scheduling may compact target rows before forward: - - dynamic_batch_size is the selected target row count + In addition: - draft_step is the candidate length to generate after target verify - pre_draft_step describes the previous iteration and controls whether GPU verify sync can be skipped """ - dynamic_batch_size: Optional[int] + origin_batch_size: int + dynamic_batch_size: int draft_step: int pre_draft_step: int # False when the current batch contains requests without a proposal from @@ -30,13 +33,9 @@ class SpecDecodePlan: # well-defined LightSpec runtime configuration. all_reqs_have_proposals: bool = True - @property - def is_dynamic(self) -> bool: - return self.dynamic_batch_size is not None - @property def skip_verify_sync(self) -> bool: - return self.is_dynamic and self.pre_draft_step == 0 + return self.pre_draft_step == 0 def filter_reqs(self, reqs: List, selected_row_mask_cpu) -> List: return [req for req, selected in zip(reqs, selected_row_mask_cpu.tolist()) if selected] @@ -46,13 +45,14 @@ class BaseMtpPlanner(ABC): """定义 SpecEngine 与不同 MTP 规划器之间的统一调用接口。""" @abstractmethod - def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: + def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: """为当前 decode 迭代生成执行计划。 Args: - decode_reqs: 当前参与 decode 的逻辑请求列表。规划器可以读取请求的 - 输出进度,判断请求是否已经持有上一轮生成的 draft proposal。 - original_batch_size: 进入动态压缩前的物理 verify 行数。 + decode_reqs: 当前参与 decode 的非空逻辑请求列表。规划器可以读取 + 请求的输出进度,判断请求是否已经持有上一轮生成的 + draft proposal。空 batch 只由 DP 专用 planner 处理。 + origin_batch_size: 进入动态压缩前的物理 verify 行数。 Returns: 本轮 target verify 使用的动态 batch size、下一轮需要生成的 draft diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index 4b28484a2f..f5501f7a8d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -28,16 +28,9 @@ def __init__(self, backend) -> None: self._register_cuda_graph_costs() self._pending_verify_batch_sizes = deque(maxlen=2) - def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: + def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: req_num = len(decode_reqs) - if req_num == 0: - return SpecDecodePlan( - dynamic_batch_size=0, - draft_step=self.max_draft_step, - pre_draft_step=self.max_draft_step, - ) - - full_batch_size = original_batch_size + full_batch_size = origin_batch_size dynamic_batch_size = full_batch_size delayed_batch_size = self._pop_delayed_batch_size( req_num=req_num, @@ -47,6 +40,7 @@ def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: dynamic_batch_size = delayed_batch_size return SpecDecodePlan( + origin_batch_size=origin_batch_size, dynamic_batch_size=dynamic_batch_size, draft_step=self.max_draft_step, pre_draft_step=self.max_draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py index fe55b6d3a6..d5db89e490 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py @@ -11,9 +11,10 @@ class FixedSpecPlanner(BaseMtpPlanner): def __init__(self, max_draft_step: int) -> None: self.max_draft_step = int(max_draft_step) - def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: + def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: return SpecDecodePlan( - dynamic_batch_size=None, + origin_batch_size=origin_batch_size, + dynamic_batch_size=origin_batch_size, draft_step=self.max_draft_step, pre_draft_step=self.max_draft_step, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index beeb942c0d..858f3ce8a0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -103,23 +103,16 @@ def _register_cuda_graph_costs(self) -> None: for batch_size, infer_cost_ms in draft_graph.infer_cost_ms_by_batch_size.items(): self.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: + def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: req_num = len(decode_reqs) pre_draft_step = self.pre_draft_step - if req_num == 0: - self.pre_draft_step = self.max_draft_step - return SpecDecodePlan( - dynamic_batch_size=0, - draft_step=self.max_draft_step, - pre_draft_step=pre_draft_step, - ) # A request entering its first decode has only the guaranteed target # row. Existing requests additionally expose pre_draft_step candidates. # This bounds Verify by the candidate pool that physically exists. req_num_with_proposals = sum(req.cur_output_len > 1 for req in decode_reqs) available_batch_size = req_num + req_num_with_proposals * pre_draft_step - max_batch_size = min(original_batch_size, available_batch_size) + max_batch_size = min(origin_batch_size, available_batch_size) all_reqs_have_proposals = req_num_with_proposals == req_num if not self.progress_ema_by_config: @@ -128,6 +121,7 @@ def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: # J(N, B, d) has a progress observation. self.pre_draft_step = self.max_draft_step return SpecDecodePlan( + origin_batch_size=origin_batch_size, dynamic_batch_size=max_batch_size, draft_step=self.max_draft_step, pre_draft_step=pre_draft_step, @@ -149,6 +143,7 @@ def plan(self, decode_reqs: List, original_batch_size: int) -> SpecDecodePlan: self.pre_draft_step = draft_step return SpecDecodePlan( + origin_batch_size=origin_batch_size, dynamic_batch_size=dynamic_batch_size, draft_step=draft_step, pre_draft_step=pre_draft_step, diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 23462f40f1..3747dd059e 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -107,15 +107,22 @@ def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): def test_fixed_planner_returns_static_plan(): planner = FixedSpecPlanner(max_draft_step=3) - plan = planner.plan(decode_reqs=build_decode_reqs(4), original_batch_size=16) + plan = planner.plan(decode_reqs=build_decode_reqs(4), origin_batch_size=16) assert isinstance(planner, BaseMtpPlanner) - assert not plan.is_dynamic - assert plan.dynamic_batch_size is None + assert plan.origin_batch_size == 16 + assert plan.dynamic_batch_size == plan.origin_batch_size assert plan.draft_step == plan.pre_draft_step == 3 assert not plan.skip_verify_sync +def test_non_dp_engine_rejects_empty_decode_batch(): + engine = SpecEngine.__new__(SpecEngine) + + with pytest.raises(AssertionError, match="requires at least one request"): + engine.plan_decode(model_input=SimpleNamespace(batch_size=0), decode_reqs=[]) + + def test_dp_planner_returns_fixed_backend_draft_step(): planner = build_dp_planner(backend=SimpleNamespace(max_draft_step=4)) @@ -288,7 +295,7 @@ def test_eagle_proposer_skips_draft_forward_for_zero_steps(): def test_dynamic_plan_filters_selected_rows(): - plan = SpecDecodePlan(dynamic_batch_size=2, draft_step=3, pre_draft_step=3) + plan = SpecDecodePlan(origin_batch_size=4, dynamic_batch_size=2, draft_step=3, pre_draft_step=3) reqs = ["req0", "req0", "req1", "req1"] selected_reqs = plan.filter_reqs( @@ -363,9 +370,10 @@ def update_mtp_verify_step_num(self, verify_step_num: int): def test_lightspec_stays_full_width_until_costs_are_profiled(): plan = build_lightspec_planner().plan( decode_reqs=build_decode_reqs(2), - original_batch_size=8, + origin_batch_size=8, ) + assert plan.origin_batch_size == 8 assert plan.dynamic_batch_size == 8 assert plan.draft_step == plan.pre_draft_step == 3 @@ -376,7 +384,7 @@ def test_lightspec_collects_full_width_progress_before_adapting(): planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) - plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) + plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) assert plan.dynamic_batch_size == 8 assert plan.draft_step == plan.pre_draft_step == 3 @@ -396,7 +404,7 @@ def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): verified_draft_step=3, ) - plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) + plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) assert plan.dynamic_batch_size == 4 assert plan.pre_draft_step == 3 @@ -420,7 +428,7 @@ def test_lightspec_compacts_block_verify_without_changing_draft_shape(): verified_draft_step=7, ) - plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=16) + plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=16) assert plan.dynamic_batch_size < 16 assert plan.draft_step == plan.pre_draft_step == 7 @@ -429,9 +437,9 @@ def test_lightspec_compacts_block_verify_without_changing_draft_shape(): def test_lightspec_bounds_verify_to_existing_proposals(): planner = build_lightspec_planner() - cold_start = planner.plan(decode_reqs=build_decode_reqs(2, 0), original_batch_size=8) - mixed_batch = planner.plan(decode_reqs=build_decode_reqs(2, 1), original_batch_size=8) - ready_batch = planner.plan(decode_reqs=build_decode_reqs(2, 2), original_batch_size=8) + cold_start = planner.plan(decode_reqs=build_decode_reqs(2, 0), origin_batch_size=8) + mixed_batch = planner.plan(decode_reqs=build_decode_reqs(2, 1), origin_batch_size=8) + ready_batch = planner.plan(decode_reqs=build_decode_reqs(2, 2), origin_batch_size=8) assert cold_start.dynamic_batch_size == 2 assert mixed_batch.dynamic_batch_size == 5 @@ -473,7 +481,7 @@ def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): verified_draft_step=3, ) - plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) + plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) assert planner.draft_steps == (1, 2, 3) assert plan.draft_step >= 1 @@ -643,7 +651,7 @@ def test_lightspec_short_current_proposal_can_recover_to_a_deeper_draft(): ) planner.pre_draft_step = 1 - plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) + plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) assert plan.dynamic_batch_size <= 4 assert plan.pre_draft_step == 1 @@ -652,9 +660,9 @@ def test_lightspec_short_current_proposal_can_recover_to_a_deeper_draft(): def test_engine_skips_feedback_for_a_mixed_proposal_batch(): engine = SpecEngine.__new__(SpecEngine) - engine.enable_dynmaic_mtp = True engine.planner = build_lightspec_planner() plan = SpecDecodePlan( + origin_batch_size=8, dynamic_batch_size=5, draft_step=3, pre_draft_step=3, @@ -680,14 +688,13 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) planner.draft_infer_costs.update(batch_size=6, infer_cost_ms=0.5) confidence_probs = np.asarray([[0.9, 0.9, 0.9]] * 2, dtype=np.float64) - plan = SpecDecodePlan(dynamic_batch_size=8, draft_step=3, pre_draft_step=3) + plan = SpecDecodePlan(origin_batch_size=8, dynamic_batch_size=8, draft_step=3, pre_draft_step=3) proposal = SpecProposal( token_ids=torch.empty((0,), dtype=torch.int64), extra_mem_indexes_cpu=None, schedule_scores_cpu=torch.from_numpy(confidence_probs), ) engine = SpecEngine.__new__(SpecEngine) - engine.enable_dynmaic_mtp = True engine.planner = planner engine.update_planner_feedback( @@ -696,14 +703,14 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): req_num=2, accept_lengths_cpu=torch.tensor([1, 1], dtype=torch.int32), ) - first_plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) + first_plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) engine.update_planner_feedback( plan=plan, proposal=proposal, req_num=2, accept_lengths_cpu=torch.tensor([1, 1], dtype=torch.int32), ) - second_plan = planner.plan(decode_reqs=build_decode_reqs(2), original_batch_size=8) + second_plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) assert first_plan.dynamic_batch_size == 8 assert second_plan.dynamic_batch_size == 4 From 4113a46d66943da2dad94f560d758102642071e7 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 02:39:45 +0000 Subject: [PATCH 047/103] refactor: extract shared MTP decode utilities --- .../mode_backend/chunked_prefill/impl.py | 19 ++- .../mode_backend/dp_backend/impl.py | 27 +++- .../model_infer/mtp_speculative/dp_engine.py | 74 +-------- .../model_infer/mtp_speculative/engine.py | 148 ++---------------- .../model_infer/mtp_speculative/utils.py | 141 +++++++++++++++++ .../test_dp_overlap_spec_engine.py | 31 ++-- .../mtp_speculative/test_planner.py | 40 +++-- 7 files changed, 235 insertions(+), 245 deletions(-) create mode 100644 lightllm/server/router/model_infer/mtp_speculative/utils.py diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 928b97bdc6..1d73d97548 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -13,6 +13,7 @@ from lightllm.server.router.model_infer.infer_batch import g_infer_context from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_current_device_id from .control_state import ControlState @@ -27,13 +28,11 @@ def __init__(self) -> None: # 用于控制每一步是执行prefill 和 decode 还是跳过 self.control_state_machine = ControlState() - self.enable_dynmaic_mtp = False # 在 mtp 模式下切换绑定的prefill 和 decode 函数 if get_env_start_args().mtp_mode is not None: self.prefill = self.prefill_mtp self.decode = self.decode_mtp - self.enable_dynmaic_mtp = get_env_start_args().mtp_dynamic_verify else: self.prefill = self.prefill_normal self.decode = self.decode_normal @@ -277,7 +276,8 @@ def decode_mtp( next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids) b_req_mtp_start_loc = gen_b_req_mtp_start_loc(model_input.b_mtp_index, num_reqs=req_num) - mtp_accept_len, accepted_index = spec_engine.verify_tokens( + mtp_accept_len, accepted_index = mtp_utils.verify_mtp_tokens( + backend=self, next_token_ids=next_token_ids, b_req_idx=model_input.b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, @@ -303,7 +303,8 @@ def decode_mtp( draft_step=spec_plan.draft_step, accept_len=mtp_accept_len, ) - spec_engine.scatter_next_tokens( + mtp_utils.scatter_mtp_next_tokens( + backend=self, b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=proposal.token_ids, b_req_idx=model_input.b_req_idx, @@ -333,7 +334,7 @@ def decode_mtp( # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() - run_reqs, verify_ok_reqs = spec_engine.resolve_decode_reqs( + run_reqs, verify_ok_reqs = mtp_utils.resolve_mtp_decode_reqs( plan=spec_plan, verify_event=verify_event, run_reqs=run_reqs, @@ -354,10 +355,11 @@ def decode_mtp( accept_lengths_cpu=mtp_accept_len_cpu, ) - spec_engine.record_request_spec_metrics( + mtp_utils.record_request_mtp_metrics( + backend=self, decode_reqs=decode_reqs, accept_lengths_cpu=mtp_accept_len_cpu, - verified_row_reqs=run_reqs if self.enable_dynmaic_mtp else None, + verified_row_reqs=run_reqs, ) select_mask = accepted_index_cpu.to(dtype=torch.bool) @@ -370,7 +372,8 @@ def decode_mtp( extra_post_req_handle_func=self.extra_post_req_handle_func, ) - spec_engine.free_unused_decode_mem( + mtp_utils.free_unused_mtp_decode_mem( + backend=self, model_input=model_input, selected_row_mask_cpu=( async_selected_row_mask_cpu.tensor if async_selected_row_mask_cpu is not None else None diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index e7ef122549..636f853acd 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -17,6 +17,7 @@ from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from .control_state import DPControlState @@ -546,7 +547,8 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): dtype=torch.int32, ).cuda(non_blocking=True) - mtp_accept_len, accepted_index = self.spec_engine.verify_tokens( + mtp_accept_len, accepted_index = mtp_utils.verify_mtp_tokens( + backend=self, next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, @@ -586,9 +588,11 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() - self.spec_engine.record_request_spec_metrics( + mtp_utils.record_request_mtp_metrics( + backend=self, decode_reqs=decode_reqs, accept_lengths_cpu=mtp_accept_len_cpu, + verified_row_reqs=run_reqs, ) verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) @@ -643,7 +647,8 @@ def _draft_decode_vanilla( ) if req_num > 0: - self.spec_engine.scatter_next_tokens( + mtp_utils.scatter_mtp_next_tokens( + backend=self, all_next_token_ids=proposal.token_ids[:req_num], b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx[:req_num], @@ -698,7 +703,8 @@ def _draft_decode_eagle( ) if req_num > 0: - self.spec_engine.scatter_next_tokens( + mtp_utils.scatter_mtp_next_tokens( + backend=self, b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=proposal.token_ids[:req_num], b_req_idx=model_input.b_req_idx[:req_num], @@ -857,7 +863,8 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if self.is_linear_att_mixed_model else None ) - mtp_accept_len, accepted_index = self.spec_engine.verify_tokens( + mtp_accept_len, accepted_index = mtp_utils.verify_mtp_tokens( + backend=self, next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, @@ -901,9 +908,11 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf if req_num0 + req_num1 > 0: event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() - self.spec_engine.record_request_spec_metrics( + mtp_utils.record_request_mtp_metrics( + backend=self, decode_reqs=decode_reqs, accept_lengths_cpu=mtp_accept_len_cpu, + verified_row_reqs=run_reqs, ) verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) @@ -977,7 +986,8 @@ def _draft_decode_vanilla_overlap( ) if req_num0 + req_num1 > 0: - self.spec_engine.scatter_next_tokens( + mtp_utils.scatter_mtp_next_tokens( + backend=self, all_next_token_ids=proposal.token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, @@ -1049,7 +1059,8 @@ def _draft_decode_eagle_overlap( ) if req_num0 + req_num1 > 0: - self.spec_engine.scatter_next_tokens( + mtp_utils.scatter_mtp_next_tokens( + backend=self, b_req_mtp_start_loc=b_req_mtp_start_loc, all_next_token_ids=proposal.token_ids, b_req_idx=b_req_idx, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py index 2d5dce8eee..cc5c0e9b43 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py @@ -1,13 +1,8 @@ -from typing import List, Optional, Tuple +from typing import Optional import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.basemodel.triton_kernel.mtp_utils import ( - linear_att_mtp_state_index_update, - mtp_scatter_next_token_ids, - mtp_verify, -) from lightllm.server.router.model_infer.mtp_speculative.dp_planner import BaseDpPlanner, build_dp_planner from lightllm.server.router.model_infer.mtp_speculative.dp_proposers import build_dp_spec_proposer from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer @@ -18,7 +13,6 @@ class DPSpecEngine: """普通 DP prefill/decode 使用的 MTP engine。""" def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: - self.backend = backend self.proposer: BaseDpProposer = build_dp_spec_proposer( spec_mode=spec_mode, backend=backend, @@ -38,31 +32,6 @@ def build_draft_state_from_prefill( next_token_ids=next_token_ids, ) - def verify_tokens( - self, - next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - b_mtp_index: Optional[torch.Tensor] = None, - ) -> Tuple[torch.Tensor, torch.Tensor]: - accept_lengths, accepted_index = mtp_verify( - req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - new_next_token_ids=next_token_ids, - b_req_idx=b_req_idx, - ) - if self.backend.is_linear_att_mixed_model: - assert b_mtp_index is not None - linear_att_mtp_state_index_update( - req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=b_req_idx, - b_mtp_index=b_mtp_index, - accepted_index=accepted_index, - verify_width=self.backend.max_draft_step + 1, - ) - return accept_lengths, accepted_index - def propose_next( self, main_model_input: ModelInput, @@ -80,46 +49,5 @@ def propose_next( accept_len=accept_len, ) - def scatter_next_tokens( - self, - b_req_mtp_start_loc: torch.Tensor, - all_next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - mtp_accept_len: torch.Tensor, - schedule_scores: Optional[torch.Tensor] = None, - ) -> None: - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, - req_to_next_token_scores=( - self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_scores - if schedule_scores is not None - else None - ), - schedule_scores=schedule_scores, - ) - - def record_request_spec_metrics( - self, - decode_reqs: List, - accept_lengths_cpu: torch.Tensor, - ) -> None: - """累计 DP 固定 verify layout 下每个请求的 MTP 指标。""" - - if not self.backend.is_master_in_dp: - return - - accept_lengths = accept_lengths_cpu.tolist() - assert len(accept_lengths) == len(decode_reqs) - for req, accept_len in zip(decode_reqs, accept_lengths): - req.update_mtp_accepted_token_num(accept_token_num=accept_len - 1) - verify_token_num = req.mtp_step + 1 - if verify_token_num > 0: - req.update_mtp_verify_token_num(verify_token_num=verify_token_num) - req.update_mtp_verify_step_num(verify_step_num=1) - __all__ = ["DPSpecEngine"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index ebf080ccd7..769d9d8d85 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -1,16 +1,10 @@ from __future__ import annotations -from collections import Counter from typing import List, Optional, Tuple import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.basemodel.triton_kernel.mtp_utils import ( - linear_att_mtp_state_index_update, - mtp_scatter_next_token_ids, - mtp_verify, -) from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager from lightllm.server.router.model_infer.mtp_speculative.planner import ( BaseMtpPlanner, @@ -24,23 +18,23 @@ class SpecEngine: - """Owns speculative planning, verification, and proposal generation. + """Owns non-DP MTP planning and draft proposal generation. - The model backend controls target forward, sampling, stream synchronization, - and request post-processing. Each proposer owns its algorithm-specific draft - state and proposal generation. + Target verification, request metrics, stream synchronization, and resource + cleanup are stateless operations exposed by ``mtp_speculative.utils``. """ def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: self.backend = backend - self.spec_mode = spec_mode - self.enable_dynmaic_mtp = enable_dynmaic_mtp self.proposer: BaseSpecProposer = build_spec_proposer( spec_mode=spec_mode, backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp, ) - self.planner: BaseMtpPlanner = self._build_decode_planner() + self.planner: BaseMtpPlanner = self._build_mtp_planner( + spec_mode=spec_mode, + enable_dynmaic_mtp=enable_dynmaic_mtp, + ) # Prefill draft-state initialization. @@ -56,7 +50,7 @@ def build_draft_state_from_prefill( next_token_ids=next_token_ids, ) - # Decode planning and target verification. + # Decode planning. def plan_decode(self, model_input: ModelInput, decode_reqs: List) -> SpecDecodePlan: """Return the fixed or dynamic speculative plan for one decode iteration.""" @@ -96,49 +90,7 @@ def prepare_decode_model_input( ) return model_input, selected_row_mask_cpu - def verify_tokens( - self, - next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - b_mtp_index: Optional[torch.Tensor] = None, - ) -> Tuple[torch.Tensor, torch.Tensor]: - accept_lengths, accepted_index = mtp_verify( - req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - new_next_token_ids=next_token_ids, - b_req_idx=b_req_idx, - ) - if self.backend.is_linear_att_mixed_model: - assert b_mtp_index is not None - linear_att_mtp_state_index_update( - req_to_mtp_state_index=self.backend.model.req_manager.req_to_mtp_state_index, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=b_req_idx, - b_mtp_index=b_mtp_index, - accepted_index=accepted_index, - verify_width=self.backend.max_draft_step + 1, - ) - return accept_lengths, accepted_index - - def resolve_decode_reqs( - self, - plan: SpecDecodePlan, - verify_event: torch.cuda.Event, - run_reqs: List, - decode_reqs: List, - accepted_index_cpu: torch.Tensor, - ) -> Tuple[List, List]: - """Return requests ready for pre/post handling after verification.""" - - if plan.skip_verify_sync: - return decode_reqs, decode_reqs - - verify_event.synchronize() - verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] - return run_reqs, verify_ok_reqs - - # Draft proposal generation and persistence. + # Draft proposal generation. def propose_next( self, @@ -158,29 +110,7 @@ def propose_next( accept_len=accept_len, ) - def scatter_next_tokens( - self, - b_req_mtp_start_loc: torch.Tensor, - all_next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - mtp_accept_len: torch.Tensor, - schedule_scores: Optional[torch.Tensor] = None, - ) -> None: - mtp_scatter_next_token_ids( - req_to_next_token_ids=self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, - b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, - req_to_next_token_scores=( - self.backend.model.req_manager.req_sampling_params_manager.req_to_next_token_scores - if schedule_scores is not None - else None - ), - schedule_scores=schedule_scores, - ) - - # Planner feedback, request metrics, and resource cleanup. + # Planner feedback. def update_planner_feedback( self, @@ -198,59 +128,9 @@ def update_planner_feedback( schedule_scores=proposal.schedule_scores_cpu, ) - def record_request_spec_metrics( - self, - decode_reqs: List, - accept_lengths_cpu: torch.Tensor, - verified_row_reqs: Optional[List] = None, - ) -> None: - """Accumulate user-visible speculative metrics on each request.""" - - if not self.backend.is_master_in_dp: - return - - accept_lengths = accept_lengths_cpu.tolist() - assert len(accept_lengths) == len(decode_reqs) - verify_rows_by_req = None if verified_row_reqs is None else Counter(req.req_idx for req in verified_row_reqs) - for req, accept_len in zip(decode_reqs, accept_lengths): - req.update_mtp_accepted_token_num(accept_token_num=accept_len - 1) - verify_token_num = req.mtp_step + 1 if verify_rows_by_req is None else verify_rows_by_req[req.req_idx] - if verify_token_num > 0: - req.update_mtp_verify_token_num(verify_token_num=verify_token_num) - req.update_mtp_verify_step_num(verify_step_num=1) - - def free_unused_decode_mem( - self, - model_input: ModelInput, - selected_row_mask_cpu: Optional[torch.Tensor], - accepted_index_cpu: torch.Tensor, - extra_mem_indexes_cpu: Optional[torch.Tensor], - ) -> None: - """Free rejected target KV slots and draft-only temporary slots.""" - - mem_indexes_cpu = model_input.mem_indexes_cpu - if selected_row_mask_cpu is None: - free_mask = accepted_index_cpu == 0 - else: - selected_mask = selected_row_mask_cpu.to(dtype=torch.bool) - free_mask = selected_mask.logical_not() - free_mask[selected_mask] = accepted_index_cpu == 0 - need_free_mem_indexes = mem_indexes_cpu[free_mask] - - if extra_mem_indexes_cpu is not None: - need_free_mem_indexes = torch.cat([need_free_mem_indexes, extra_mem_indexes_cpu], dim=0) - if len(need_free_mem_indexes) > 0: - self.backend.model.req_manager.mem_manager.free(need_free_mem_indexes) - - # Construction helpers. - - def _build_decode_planner(self) -> BaseMtpPlanner: - if not self.enable_dynmaic_mtp: + def _build_mtp_planner(self, spec_mode: str, enable_dynmaic_mtp: bool) -> BaseMtpPlanner: + if not enable_dynmaic_mtp: return FixedSpecPlanner(max_draft_step=self.backend.max_draft_step) - if self.spec_mode == "dspark": + if spec_mode == "dspark": return DSparkPlanner(backend=self.backend) - - return LightSpecPlanner( - spec_mode=self.spec_mode, - backend=self.backend, - ) + return LightSpecPlanner(spec_mode=spec_mode, backend=self.backend) diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py new file mode 100644 index 0000000000..a0f947bc5b --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +from collections import Counter +from typing import TYPE_CHECKING, List, Optional, Tuple + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput +from lightllm.common.basemodel.triton_kernel.mtp_utils import ( + linear_att_mtp_state_index_update, + mtp_scatter_next_token_ids, + mtp_verify, +) +from lightllm.server.router.model_infer.mtp_speculative.planner import SpecDecodePlan + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.infer_batch import InferReq + + +def verify_mtp_tokens( + backend, + next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + b_mtp_index: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Verify target tokens and update recurrent MTP state when required.""" + + accept_lengths, accepted_index = mtp_verify( + req_to_next_token_ids=backend.model.req_manager.req_sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + new_next_token_ids=next_token_ids, + b_req_idx=b_req_idx, + ) + if backend.is_linear_att_mixed_model: + assert b_mtp_index is not None + linear_att_mtp_state_index_update( + req_to_mtp_state_index=backend.model.req_manager.req_to_mtp_state_index, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_req_idx=b_req_idx, + b_mtp_index=b_mtp_index, + accepted_index=accepted_index, + verify_width=backend.max_draft_step + 1, + ) + return accept_lengths, accepted_index + + +def resolve_mtp_decode_reqs( + plan: SpecDecodePlan, + verify_event: torch.cuda.Event, + run_reqs: List, + decode_reqs: List, + accepted_index_cpu: torch.Tensor, +) -> Tuple[List, List]: + """Return requests ready for pre/post handling after verification.""" + + if plan.skip_verify_sync: + return decode_reqs, decode_reqs + + verify_event.synchronize() + verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] + return run_reqs, verify_ok_reqs + + +def scatter_mtp_next_tokens( + backend, + b_req_mtp_start_loc: torch.Tensor, + all_next_token_ids: torch.Tensor, + b_req_idx: torch.Tensor, + mtp_accept_len: torch.Tensor, + schedule_scores: Optional[torch.Tensor] = None, +) -> None: + """Persist the next MTP proposal and optional scheduling scores by request.""" + + sampling_params_manager = backend.model.req_manager.req_sampling_params_manager + mtp_scatter_next_token_ids( + req_to_next_token_ids=sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + all_next_token_ids=all_next_token_ids, + b_req_idx=b_req_idx, + mtp_accept_len=mtp_accept_len, + req_to_next_token_scores=( + sampling_params_manager.req_to_next_token_scores if schedule_scores is not None else None + ), + schedule_scores=schedule_scores, + ) + + +def record_request_mtp_metrics( + backend, + decode_reqs: List[InferReq], + accept_lengths_cpu: torch.Tensor, + verified_row_reqs: List[InferReq], +) -> None: + """Accumulate user-visible MTP metrics on each request.""" + + if not backend.is_master_in_dp: + return + + accept_lengths = accept_lengths_cpu.tolist() + assert len(accept_lengths) == len(decode_reqs) + verify_rows_by_req = Counter(req.req_idx for req in verified_row_reqs) + for req, accept_len in zip(decode_reqs, accept_lengths): + req.update_mtp_accepted_token_num(accept_token_num=accept_len - 1) + verify_token_num = verify_rows_by_req[req.req_idx] + if verify_token_num > 0: + req.update_mtp_verify_token_num(verify_token_num=verify_token_num) + req.update_mtp_verify_step_num(verify_step_num=1) + + +def free_unused_mtp_decode_mem( + backend, + model_input: ModelInput, + selected_row_mask_cpu: Optional[torch.Tensor], + accepted_index_cpu: torch.Tensor, + extra_mem_indexes_cpu: Optional[torch.Tensor], +) -> None: + """Free rejected target KV slots and draft-only temporary slots.""" + + mem_indexes_cpu = model_input.mem_indexes_cpu + if selected_row_mask_cpu is None: + free_mask = accepted_index_cpu == 0 + else: + selected_mask = selected_row_mask_cpu.to(dtype=torch.bool) + free_mask = selected_mask.logical_not() + free_mask[selected_mask] = accepted_index_cpu == 0 + need_free_mem_indexes = mem_indexes_cpu[free_mask] + + if extra_mem_indexes_cpu is not None: + need_free_mem_indexes = torch.cat([need_free_mem_indexes, extra_mem_indexes_cpu], dim=0) + if len(need_free_mem_indexes) > 0: + backend.model.req_manager.mem_manager.free(need_free_mem_indexes) + + +__all__ = [ + "free_unused_mtp_decode_mem", + "record_request_mtp_metrics", + "resolve_mtp_decode_reqs", + "scatter_mtp_next_tokens", + "verify_mtp_tokens", +] diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 14e2e82620..1449ed7162 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -2,6 +2,7 @@ import torch +from lightllm.server.router.model_infer.mode_backend.dp_backend import impl as dp_backend_impl from lightllm.server.router.model_infer.mode_backend.dp_backend.impl import DPChunkedPrefillBackend from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ChunkedPrefillBackend @@ -15,7 +16,6 @@ class _RecordingSpecEngine: def __init__(self): self.propose_args = None self.propose_overlap_args = None - self.scatter_args = None def propose_next(self, **kwargs): self.propose_args = kwargs @@ -34,8 +34,13 @@ def propose_next_overlap(self, **kwargs): extra_mem_indexes_cpu=torch.tensor([456], dtype=torch.int32), ) - def scatter_next_tokens(self, **kwargs): - self.scatter_args = kwargs + +def _capture_scatter_args(monkeypatch): + scatter_args = {} + monkeypatch.setattr( + dp_backend_impl.mtp_utils, "scatter_mtp_next_tokens", lambda **kwargs: scatter_args.update(kwargs) + ) + return scatter_args def test_backends_initialize_their_own_spec_engine(): @@ -90,7 +95,8 @@ def test_padded_token_ids_support_empty_dp_rank(): assert torch.equal(padded_token_ids, torch.zeros(4, dtype=torch.int64)) -def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(): +def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(monkeypatch): + scatter_args = _capture_scatter_args(monkeypatch) backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) backend.max_draft_step = 7 backend.spec_engine = _RecordingSpecEngine() @@ -118,11 +124,12 @@ def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(): assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) assert torch.equal(propose_args["b_req_mtp_start_loc"], torch.tensor([0, 8], dtype=torch.int32)) assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) - assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (8, 8) + assert scatter_args["all_next_token_ids"].shape == (8, 8) assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) -def test_dp_vanilla_uses_dp_engine_proposer(): +def test_dp_vanilla_uses_dp_engine_proposer(monkeypatch): + scatter_args = _capture_scatter_args(monkeypatch) backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) backend.max_draft_step = 7 backend.spec_engine = _RecordingSpecEngine() @@ -145,11 +152,12 @@ def test_dp_vanilla_uses_dp_engine_proposer(): propose_args = backend.spec_engine.propose_args assert propose_args["next_token_ids"].shape == (16,) assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) - assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (8, 8) + assert scatter_args["all_next_token_ids"].shape == (8, 8) assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) -def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(): +def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(monkeypatch): + scatter_args = _capture_scatter_args(monkeypatch) backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) backend.max_draft_step = 7 backend.spec_engine = _RecordingSpecEngine() @@ -190,11 +198,12 @@ def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(): assert propose_args["real_verify_rows1"] == 16 assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) - assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (24, 8) + assert scatter_args["all_next_token_ids"].shape == (24, 8) assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) -def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(): +def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(monkeypatch): + scatter_args = _capture_scatter_args(monkeypatch) backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) backend.max_draft_step = 7 backend.spec_engine = _RecordingSpecEngine() @@ -230,5 +239,5 @@ def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(): assert propose_args["real_verify_rows1"] == 16 assert propose_args["accept_len0"] is None assert propose_args["accept_len1"] is None - assert backend.spec_engine.scatter_args["all_next_token_ids"].shape == (24, 8) + assert scatter_args["all_next_token_ids"].shape == (24, 8) assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 3747dd059e..6fe805da96 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -58,6 +58,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import VanillaWithAttProposer +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils def build_lightspec_planner( @@ -95,14 +96,15 @@ def build_decode_reqs(req_num: int, req_num_with_proposals: int | None = None): def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): engine = SpecEngine.__new__(SpecEngine) - engine.spec_mode = spec_mode - engine.enable_dynmaic_mtp = enable_dynmaic_mtp engine.backend = SimpleNamespace( max_draft_step=3, model=SimpleNamespace(graph=None), draft_models=[SimpleNamespace(block_size=3, graph=None)], ) - return engine._build_decode_planner() + return engine._build_mtp_planner( + spec_mode=spec_mode, + enable_dynmaic_mtp=enable_dynmaic_mtp, + ) def test_fixed_planner_returns_static_plan(): @@ -123,6 +125,20 @@ def test_non_dp_engine_rejects_empty_decode_batch(): engine.plan_decode(model_input=SimpleNamespace(batch_size=0), decode_reqs=[]) +def test_spec_engine_only_exposes_planning_and_proposal_interfaces(): + public_methods = { + name for name, value in SpecEngine.__dict__.items() if callable(value) and not name.startswith("_") + } + + assert public_methods == { + "build_draft_state_from_prefill", + "plan_decode", + "prepare_decode_model_input", + "propose_next", + "update_planner_feedback", + } + + def test_dp_planner_returns_fixed_backend_draft_step(): planner = build_dp_planner(backend=SimpleNamespace(max_draft_step=4)) @@ -308,15 +324,15 @@ def test_dynamic_plan_filters_selected_rows(): def test_dynamic_decode_frees_unselected_and_rejected_rows(): freed = [] - engine = SpecEngine.__new__(SpecEngine) - engine.backend = SimpleNamespace( + backend = SimpleNamespace( model=SimpleNamespace( req_manager=SimpleNamespace(mem_manager=SimpleNamespace(free=lambda indexes: freed.append(indexes.clone()))) ) ) model_input = SimpleNamespace(mem_indexes_cpu=torch.tensor([10, 11, 12, 13])) - engine.free_unused_decode_mem( + mtp_utils.free_unused_mtp_decode_mem( + backend=backend, model_input=model_input, selected_row_mask_cpu=torch.tensor([1, 0, 1, 0], dtype=torch.bool), accepted_index_cpu=torch.tensor([1, 0], dtype=torch.int32), @@ -327,7 +343,7 @@ def test_dynamic_decode_frees_unselected_and_rejected_rows(): assert freed[0].tolist() == [11, 12, 13, 20] -def test_engine_records_request_spec_metrics_in_one_pass(): +def test_records_request_mtp_metrics_in_one_pass(): class MetricReq: def __init__(self, req_idx: int, mtp_step: int): self.req_idx = req_idx @@ -345,12 +361,12 @@ def update_mtp_verify_token_num(self, verify_token_num: int): def update_mtp_verify_step_num(self, verify_step_num: int): self.verify_steps += verify_step_num - engine = SpecEngine.__new__(SpecEngine) - engine.backend = SimpleNamespace(is_master_in_dp=True) + backend = SimpleNamespace(is_master_in_dp=True) req0 = MetricReq(req_idx=10, mtp_step=3) req1 = MetricReq(req_idx=11, mtp_step=3) - engine.record_request_spec_metrics( + mtp_utils.record_request_mtp_metrics( + backend=backend, decode_reqs=[req0, req1], accept_lengths_cpu=torch.tensor([2, 1], dtype=torch.int32), verified_row_reqs=[req0, req0, req1], @@ -360,9 +376,11 @@ def update_mtp_verify_step_num(self, verify_step_num: int): assert (req1.accepted, req1.verified, req1.verify_steps) == (0, 1, 1) fixed_req = MetricReq(req_idx=12, mtp_step=3) - engine.record_request_spec_metrics( + mtp_utils.record_request_mtp_metrics( + backend=backend, decode_reqs=[fixed_req], accept_lengths_cpu=torch.tensor([3], dtype=torch.int32), + verified_row_reqs=[fixed_req] * 4, ) assert (fixed_req.accepted, fixed_req.verified, fixed_req.verify_steps) == (2, 4, 1) From f78238fd2a5d9804404632035201b5b3c6ccaa4e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 02:54:19 +0000 Subject: [PATCH 048/103] refactor: simplify MTP decode request handling --- .../mode_backend/chunked_prefill/impl.py | 25 +++++++++-------- .../mtp_speculative/planner/base.py | 3 -- .../model_infer/mtp_speculative/utils.py | 28 ++++--------------- .../mtp_speculative/test_planner.py | 12 -------- 4 files changed, 19 insertions(+), 49 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 1d73d97548..3b815ae060 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -262,12 +262,13 @@ def decode_mtp( ) model_output = self.model.forward(model_input) + # 动态 MTP verify 可能只从原始物理 batch 中选择部分行参与 target forward。 + # 等待异步回传的行选择掩码后,按相同掩码过滤 run_reqs,使请求列表的 + # 长度和顺序与压缩后的 model_output.logits 保持一一对应,供后续采样使用。 if async_selected_row_mask_cpu is not None: async_selected_row_mask_cpu.wait() - run_reqs = spec_plan.filter_reqs( - reqs=run_reqs, - selected_row_mask_cpu=async_selected_row_mask_cpu.tensor, - ) + selected_row_mask_cpu = async_selected_row_mask_cpu.tensor.tolist() + run_reqs = [req for req, selected in zip(run_reqs, selected_row_mask_cpu) if selected] next_token_ids, next_token_logprobs = sample( model_output.logits, run_reqs, @@ -334,13 +335,15 @@ def decode_mtp( # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() - run_reqs, verify_ok_reqs = mtp_utils.resolve_mtp_decode_reqs( - plan=spec_plan, - verify_event=verify_event, - run_reqs=run_reqs, - decode_reqs=decode_reqs, - accepted_index_cpu=accepted_index_cpu, - ) + # 当 pre_draft_step == 0 时,上一轮没有生成 draft token,本轮每个请求 + # 只有一个由 target model 产生且必然提交的 token,不存在需要根据 + # accepted_index_cpu 剔除的 draft 行。因此这里可以直接使用 run_reqs, + # 无需等待 verify_event,避免一次不必要的 GPU/CPU 同步。 + if spec_plan.skip_verify_sync: + verify_ok_reqs = run_reqs + else: + verify_event.synchronize() + verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py index 7802421632..adcc996116 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -37,9 +37,6 @@ class SpecDecodePlan: def skip_verify_sync(self) -> bool: return self.pre_draft_step == 0 - def filter_reqs(self, reqs: List, selected_row_mask_cpu) -> List: - return [req for req, selected in zip(reqs, selected_row_mask_cpu.tolist()) if selected] - class BaseMtpPlanner(ABC): """定义 SpecEngine 与不同 MTP 规划器之间的统一调用接口。""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index a0f947bc5b..5fbf15a58f 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -11,14 +11,14 @@ mtp_scatter_next_token_ids, mtp_verify, ) -from lightllm.server.router.model_infer.mtp_speculative.planner import SpecDecodePlan if TYPE_CHECKING: from lightllm.server.router.model_infer.infer_batch import InferReq + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend def verify_mtp_tokens( - backend, + backend: ModeBackend, next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, @@ -45,25 +45,8 @@ def verify_mtp_tokens( return accept_lengths, accepted_index -def resolve_mtp_decode_reqs( - plan: SpecDecodePlan, - verify_event: torch.cuda.Event, - run_reqs: List, - decode_reqs: List, - accepted_index_cpu: torch.Tensor, -) -> Tuple[List, List]: - """Return requests ready for pre/post handling after verification.""" - - if plan.skip_verify_sync: - return decode_reqs, decode_reqs - - verify_event.synchronize() - verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] - return run_reqs, verify_ok_reqs - - def scatter_mtp_next_tokens( - backend, + backend: ModeBackend, b_req_mtp_start_loc: torch.Tensor, all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, @@ -87,7 +70,7 @@ def scatter_mtp_next_tokens( def record_request_mtp_metrics( - backend, + backend: ModeBackend, decode_reqs: List[InferReq], accept_lengths_cpu: torch.Tensor, verified_row_reqs: List[InferReq], @@ -109,7 +92,7 @@ def record_request_mtp_metrics( def free_unused_mtp_decode_mem( - backend, + backend: ModeBackend, model_input: ModelInput, selected_row_mask_cpu: Optional[torch.Tensor], accepted_index_cpu: torch.Tensor, @@ -135,7 +118,6 @@ def free_unused_mtp_decode_mem( __all__ = [ "free_unused_mtp_decode_mem", "record_request_mtp_metrics", - "resolve_mtp_decode_reqs", "scatter_mtp_next_tokens", "verify_mtp_tokens", ] diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 6fe805da96..bcda6c892e 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -310,18 +310,6 @@ def test_eagle_proposer_skips_draft_forward_for_zero_steps(): assert proposal.schedule_scores.shape == (2, 0) -def test_dynamic_plan_filters_selected_rows(): - plan = SpecDecodePlan(origin_batch_size=4, dynamic_batch_size=2, draft_step=3, pre_draft_step=3) - reqs = ["req0", "req0", "req1", "req1"] - - selected_reqs = plan.filter_reqs( - reqs=reqs, - selected_row_mask_cpu=torch.tensor([1, 0, 1, 0], dtype=torch.int32), - ) - - assert selected_reqs == ["req0", "req1"] - - def test_dynamic_decode_frees_unselected_and_rejected_rows(): freed = [] backend = SimpleNamespace( From 29940cb97f9f879d5bddac17a23f2a26c0ac3850 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 03:06:27 +0000 Subject: [PATCH 049/103] refactor: clarify MTP planner interfaces --- .../mode_backend/chunked_prefill/impl.py | 2 +- .../model_infer/mtp_speculative/engine.py | 8 +- .../mtp_speculative/planner/base.py | 36 +++---- .../mtp_speculative/planner/dspark.py | 32 +++---- .../mtp_speculative/planner/fixed.py | 2 +- .../mtp_speculative/planner/lightspec.py | 94 +++++++++---------- .../mtp_speculative/test_planner.py | 40 ++++---- 7 files changed, 107 insertions(+), 107 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 3b815ae060..36699d122a 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -351,7 +351,7 @@ def decode_mtp( event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() - spec_engine.update_planner_feedback( + spec_engine.update_planner_statics( plan=spec_plan, proposal=proposal, req_num=req_num, diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 769d9d8d85..2ed0c3e0bb 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -110,18 +110,18 @@ def propose_next( accept_len=accept_len, ) - # Planner feedback. + # Planner runtime statistics. - def update_planner_feedback( + def update_planner_statics( self, plan: SpecDecodePlan, proposal: SpecProposal, req_num: int, accept_lengths_cpu: torch.Tensor, ) -> None: - """Feed iteration-level observations into the current planner.""" + """Update the current planner with iteration-level runtime statistics.""" - self.planner.update_feedback( + self.planner.update_statics( plan=plan, req_num=req_num, accept_lengths=accept_lengths_cpu, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py index adcc996116..e278157cd0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -59,7 +59,7 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: raise NotImplementedError @abstractmethod - def update_feedback( + def update_statics( self, plan: SpecDecodePlan, req_num: int, @@ -87,7 +87,23 @@ def __init__(self) -> None: def update(self, batch_size: int, infer_cost_ms: float) -> None: self.infer_cost_ms_table[int(batch_size)] = float(infer_cost_ms) - def get(self, batch_size: int) -> float: + def estimate(self, batch_size: int) -> float: + """Estimate an uncaptured batch without applying the graph-miss penalty.""" + + batch_size = int(batch_size) + max_batch_size, max_cost_ms = self.infer_cost_ms_table.peekitem(-1) + if batch_size <= max_batch_size: + return self._get(batch_size) + return max_cost_ms * batch_size / max_batch_size + + def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> List[int]: + start = min(int(batch_size1), int(batch_size2)) + end = max(int(batch_size1), int(batch_size2)) + batch_sizes = set(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) + batch_sizes.update((start, end)) + return sorted(batch_sizes) + + def _get(self, batch_size: int) -> float: batch_size = int(batch_size) if len(self.infer_cost_ms_table) == 0: @@ -105,22 +121,6 @@ def get(self, batch_size: int) -> float: index = self.infer_cost_ms_table.bisect_left(batch_size) return self.infer_cost_ms_table.peekitem(index)[1] - def estimate(self, batch_size: int) -> float: - """Estimate an uncaptured batch without applying the graph-miss penalty.""" - - batch_size = int(batch_size) - max_batch_size, max_cost_ms = self.infer_cost_ms_table.peekitem(-1) - if batch_size <= max_batch_size: - return self.get(batch_size) - return max_cost_ms * batch_size / max_batch_size - - def get_batch_size_keys_between(self, batch_size1: int, batch_size2: int) -> List[int]: - start = min(int(batch_size1), int(batch_size2)) - end = max(int(batch_size1), int(batch_size2)) - batch_sizes = set(self.infer_cost_ms_table.irange(minimum=start, maximum=end, inclusive=(True, True))) - batch_sizes.update((start, end)) - return sorted(batch_sizes) - class _EMAValue: def __init__(self, decay: float, init_value: float) -> None: diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index f5501f7a8d..cc1301af97 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -46,7 +46,20 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: pre_draft_step=self.max_draft_step, ) - def get_draft_cost_ms(self, req_num: int, verify_batch_size: int, draft_step: int) -> float: + def update_statics( + self, + plan: SpecDecodePlan, + req_num: int, + accept_lengths, + schedule_scores=None, + ) -> None: + if schedule_scores is not None: + self._update_confidence_probs( + confidence_probs=schedule_scores, + req_num=req_num, + ) + + def _get_draft_cost_ms(self, req_num: int, verify_batch_size: int, draft_step: int) -> float: """Return the cost of committing verify rows and generating one complete block.""" extend_cost_ms = self.draft_infer_costs.estimate(verify_batch_size) @@ -66,20 +79,7 @@ def _register_cuda_graph_costs(self) -> None: for batch_size, infer_cost_ms in draft_graph.infer_cost_ms_by_batch_size.items(): self.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - def update_feedback( - self, - plan: SpecDecodePlan, - req_num: int, - accept_lengths, - schedule_scores=None, - ) -> None: - if schedule_scores is not None: - self.update_confidence_probs( - confidence_probs=schedule_scores, - req_num=req_num, - ) - - def update_confidence_probs(self, confidence_probs, req_num: int) -> None: + def _update_confidence_probs(self, confidence_probs, req_num: int) -> None: """Record a confidence-derived future capacity estimate. The current decode step still routes rows by the current probabilities @@ -150,7 +150,7 @@ def _select_dynamic_batch_size_from_survival_scores( for dynamic_batch_size in sorted(candidate_batch_sizes): selected_draft_count = dynamic_batch_size - int(req_num) expected_tokens = float(req_num) + float(expected_accepts_by_count[selected_draft_count]) - round_ms = self.target_infer_costs.estimate(dynamic_batch_size) + self.get_draft_cost_ms( + round_ms = self.target_infer_costs.estimate(dynamic_batch_size) + self._get_draft_cost_ms( req_num=req_num, verify_batch_size=dynamic_batch_size, draft_step=self.max_draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py index d5db89e490..34976164db 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py @@ -19,7 +19,7 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: pre_draft_step=self.max_draft_step, ) - def update_feedback( + def update_statics( self, plan: SpecDecodePlan, req_num: int, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index 858f3ce8a0..c0c70b3172 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -43,7 +43,7 @@ def __init__( self.backend = backend self.max_draft_step = int(backend.max_draft_step) self.block_size = int(backend.draft_models[0].block_size) if self.spec_mode == "dflash" else None - self.draft_steps = self.get_draft_steps() + self.draft_steps = self._get_draft_steps() self.target_infer_costs = _InferCostMsTable() self.draft_infer_costs = _InferCostMsTable() @@ -61,48 +61,6 @@ def __init__( # The current verify width is bounded by the proposal built last time. self.pre_draft_step = self.max_draft_step - def get_draft_steps(self) -> Tuple[int, ...]: - """Return the draft configurations supported by the current MTP mode.""" - - if self.spec_mode in ("vanilla_no_att", "eagle_no_att"): - return tuple(range(self.max_draft_step + 1)) - if self.spec_mode in ("vanilla_with_att", "eagle_with_att", "eagle3"): - return tuple(range(1, self.max_draft_step + 1)) - if self.spec_mode == "dflash": - return (self.max_draft_step,) - raise ValueError(f"unsupported LightSpec mode: {self.spec_mode}") - - def get_draft_cost_ms(self, req_num: int, verify_batch_size: int, draft_step: int) -> float: - """Return the complete draft cost for one ``(N, B, d)`` configuration.""" - - if self.spec_mode in ("vanilla_no_att", "eagle_no_att"): - return self.draft_infer_costs.estimate(req_num) * draft_step - if self.spec_mode in ("vanilla_with_att", "eagle_with_att", "eagle3"): - assert draft_step > 0, f"{self.spec_mode} requires draft_step to be greater than 0" - draft_cost_ms = self.draft_infer_costs.estimate(verify_batch_size) - if draft_step > 1: - draft_cost_ms += self.draft_infer_costs.estimate(req_num) * (draft_step - 1) - return draft_cost_ms - if self.spec_mode == "dflash": - assert self.block_size is not None - extend_cost_ms = self.draft_infer_costs.estimate(verify_batch_size) - block_cost_ms = self.draft_infer_costs.estimate(req_num * self.block_size) - return extend_cost_ms + block_cost_ms - raise ValueError(f"unsupported LightSpec mode: {self.spec_mode}") - - def _register_cuda_graph_costs(self) -> None: - target_graph = self.backend.model.graph - if target_graph is not None: - for batch_size, infer_cost_ms in target_graph.infer_cost_ms_by_batch_size.items(): - self.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - - for draft_model in self.backend.draft_models: - draft_graph = draft_model.graph - if draft_graph is None: - continue - for batch_size, infer_cost_ms in draft_graph.infer_cost_ms_by_batch_size.items(): - self.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) - def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: req_num = len(decode_reqs) pre_draft_step = self.pre_draft_step @@ -150,7 +108,7 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: all_reqs_have_proposals=all_reqs_have_proposals, ) - def update_feedback( + def update_statics( self, plan: SpecDecodePlan, req_num: int, @@ -165,14 +123,56 @@ def update_feedback( # updating the batch-level statistic with different N/B semantics. if not plan.all_reqs_have_proposals: return - self.update_verified_batch( + self._update_verified_batch( accept_lengths=accept_lengths, req_num=req_num, dynamic_batch_size=plan.dynamic_batch_size, verified_draft_step=plan.pre_draft_step, ) - def update_verified_batch( + def _get_draft_steps(self) -> Tuple[int, ...]: + """Return the draft configurations supported by the current MTP mode.""" + + if self.spec_mode in ("vanilla_no_att", "eagle_no_att"): + return tuple(range(self.max_draft_step + 1)) + if self.spec_mode in ("vanilla_with_att", "eagle_with_att", "eagle3"): + return tuple(range(1, self.max_draft_step + 1)) + if self.spec_mode == "dflash": + return (self.max_draft_step,) + raise ValueError(f"unsupported LightSpec mode: {self.spec_mode}") + + def _get_draft_cost_ms(self, req_num: int, verify_batch_size: int, draft_step: int) -> float: + """Return the complete draft cost for one ``(N, B, d)`` configuration.""" + + if self.spec_mode in ("vanilla_no_att", "eagle_no_att"): + return self.draft_infer_costs.estimate(req_num) * draft_step + if self.spec_mode in ("vanilla_with_att", "eagle_with_att", "eagle3"): + assert draft_step > 0, f"{self.spec_mode} requires draft_step to be greater than 0" + draft_cost_ms = self.draft_infer_costs.estimate(verify_batch_size) + if draft_step > 1: + draft_cost_ms += self.draft_infer_costs.estimate(req_num) * (draft_step - 1) + return draft_cost_ms + if self.spec_mode == "dflash": + assert self.block_size is not None + extend_cost_ms = self.draft_infer_costs.estimate(verify_batch_size) + block_cost_ms = self.draft_infer_costs.estimate(req_num * self.block_size) + return extend_cost_ms + block_cost_ms + raise ValueError(f"unsupported LightSpec mode: {self.spec_mode}") + + def _register_cuda_graph_costs(self) -> None: + target_graph = self.backend.model.graph + if target_graph is not None: + for batch_size, infer_cost_ms in target_graph.infer_cost_ms_by_batch_size.items(): + self.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + + for draft_model in self.backend.draft_models: + draft_graph = draft_model.graph + if draft_graph is None: + continue + for batch_size, infer_cost_ms in draft_graph.infer_cost_ms_by_batch_size.items(): + self.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=infer_cost_ms) + + def _update_verified_batch( self, accept_lengths, req_num: int, @@ -244,7 +244,7 @@ def _get_cost_ms(self, req_num: int, dynamic_batch_size: int, draft_step: int) - dynamic_batch_size=dynamic_batch_size, draft_step=draft_step, ) - total_time = self.target_infer_costs.estimate(dynamic_batch_size) + self.get_draft_cost_ms( + total_time = self.target_infer_costs.estimate(dynamic_batch_size) + self._get_draft_cost_ms( req_num=req_num, verify_batch_size=dynamic_batch_size, draft_step=draft_step, diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index bcda6c892e..8a4e814d48 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -135,7 +135,7 @@ def test_spec_engine_only_exposes_planning_and_proposal_interfaces(): "plan_decode", "prepare_decode_model_input", "propose_next", - "update_planner_feedback", + "update_planner_statics", } @@ -403,7 +403,7 @@ def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): for batch_size, draft_cost in ((2, 0.1), (4, 0.1), (8, 0.2)): planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=draft_cost) for _ in range(8): - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[4, 4], req_num=2, dynamic_batch_size=8, @@ -427,7 +427,7 @@ def test_lightspec_compacts_block_verify_without_changing_draft_shape(): planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) planner.draft_infer_costs.update(batch_size=14, infer_cost_ms=0.5) for _ in range(8): - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[3, 3], req_num=2, dynamic_batch_size=16, @@ -480,7 +480,7 @@ def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): for batch_size in (2, 4, 8): planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[2, 2], req_num=2, dynamic_batch_size=8, @@ -497,7 +497,7 @@ def test_vanilla_with_attention_planner_prices_extend_then_normal_batches(): planner = build_lightspec_planner(spec_mode="vanilla_with_att") planner.draft_infer_costs.update(batch_size=2, infer_cost_ms=0.1) - draft_cost_ms = planner.get_draft_cost_ms( + draft_cost_ms = planner._get_draft_cost_ms( req_num=4, verify_batch_size=8, draft_step=3, @@ -505,7 +505,7 @@ def test_vanilla_with_attention_planner_prices_extend_then_normal_batches(): assert np.isclose(draft_cost_ms, 0.8) with pytest.raises(AssertionError, match="requires draft_step to be greater than 0"): - planner.get_draft_cost_ms(req_num=4, verify_batch_size=8, draft_step=0) + planner._get_draft_cost_ms(req_num=4, verify_batch_size=8, draft_step=0) def test_no_attention_planner_prices_only_normal_request_batches(): @@ -513,7 +513,7 @@ def test_no_attention_planner_prices_only_normal_request_batches(): planner = build_lightspec_planner(spec_mode=spec_mode) planner.draft_infer_costs.update(batch_size=2, infer_cost_ms=0.1) - draft_cost_ms = planner.get_draft_cost_ms( + draft_cost_ms = planner._get_draft_cost_ms( req_num=4, verify_batch_size=8, draft_step=3, @@ -526,7 +526,7 @@ def test_eagle_planner_prices_extend_and_decode_separately(): planner = build_lightspec_planner(spec_mode="eagle_with_att") planner.draft_infer_costs.update(batch_size=2, infer_cost_ms=0.25) - draft_cost_ms = planner.get_draft_cost_ms( + draft_cost_ms = planner._get_draft_cost_ms( req_num=2, verify_batch_size=8, draft_step=3, @@ -544,7 +544,7 @@ def test_autoregressive_eagle_planner_prices_extend_and_decode_rows(): for batch_size in (2, 4, 8, 16): planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=float(batch_size)) - draft_cost_ms = planner.get_draft_cost_ms( + draft_cost_ms = planner._get_draft_cost_ms( req_num=8, verify_batch_size=16, draft_step=7, @@ -561,7 +561,7 @@ def test_block_planner_prices_commit_and_complete_block(): ) planner.draft_infer_costs.update(batch_size=8, infer_cost_ms=0.4) - draft_cost_ms = planner.get_draft_cost_ms( + draft_cost_ms = planner._get_draft_cost_ms( req_num=2, verify_batch_size=16, draft_step=7, @@ -574,7 +574,7 @@ def test_dspark_planner_prices_commit_and_complete_block(): planner = build_dspark_planner(max_draft_step=3, block_size=3) planner.draft_infer_costs.update(batch_size=8, infer_cost_ms=0.8) - draft_cost_ms = planner.get_draft_cost_ms( + draft_cost_ms = planner._get_draft_cost_ms( req_num=4, verify_batch_size=8, draft_step=3, @@ -586,13 +586,13 @@ def test_dspark_planner_prices_commit_and_complete_block(): def test_lightspec_records_one_batch_observation_per_configuration(): planner = build_lightspec_planner() - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[2, 2], req_num=2, dynamic_batch_size=4, verified_draft_step=1, ) - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[1, 1], req_num=2, dynamic_batch_size=4, @@ -608,7 +608,7 @@ def test_lightspec_records_one_batch_observation_per_configuration(): def test_lightspec_high_concurrency_does_not_multiply_ema_updates(): planner = build_lightspec_planner() - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[1] * 128, req_num=128, dynamic_batch_size=128, @@ -620,7 +620,7 @@ def test_lightspec_high_concurrency_does_not_multiply_ema_updates(): def test_lightspec_estimates_unseen_shapes_from_prefix_survival(): planner = build_lightspec_planner() - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[2, 1], req_num=2, dynamic_batch_size=8, @@ -634,7 +634,7 @@ def test_lightspec_estimates_unseen_shapes_from_prefix_survival(): def test_lightspec_does_not_transfer_deep_progress_to_short_drafts(): planner = build_lightspec_planner() - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[4, 2], req_num=2, dynamic_batch_size=8, @@ -649,7 +649,7 @@ def test_lightspec_short_current_proposal_can_recover_to_a_deeper_draft(): for batch_size, target_cost in ((2, 1.0), (4, 1.1), (6, 1.2), (8, 1.3)): planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=0.01) - planner.update_verified_batch( + planner._update_verified_batch( accept_lengths=[4, 4], req_num=2, dynamic_batch_size=8, @@ -675,7 +675,7 @@ def test_engine_skips_feedback_for_a_mixed_proposal_batch(): all_reqs_have_proposals=False, ) - engine.update_planner_feedback( + engine.update_planner_statics( plan=plan, proposal=SpecProposal( token_ids=torch.empty((0,), dtype=torch.int64), @@ -703,14 +703,14 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): engine = SpecEngine.__new__(SpecEngine) engine.planner = planner - engine.update_planner_feedback( + engine.update_planner_statics( plan=plan, proposal=proposal, req_num=2, accept_lengths_cpu=torch.tensor([1, 1], dtype=torch.int32), ) first_plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) - engine.update_planner_feedback( + engine.update_planner_statics( plan=plan, proposal=proposal, req_num=2, From c5c941a65a9cc693750c9ec704267fdc6e627236 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 03:23:26 +0000 Subject: [PATCH 050/103] refactor: type MTP backend dependencies --- lightllm/server/router/model_infer/infer_batch.py | 9 +++++---- .../router/model_infer/mtp_speculative/dp_engine.py | 9 +++++++-- .../model_infer/mtp_speculative/dp_overlap_engine.py | 9 +++++++-- .../mtp_speculative/dp_overlap_planner/__init__.py | 7 ++++++- .../mtp_speculative/dp_overlap_proposers/__init__.py | 3 ++- .../model_infer/mtp_speculative/dp_planner/__init__.py | 7 ++++++- .../model_infer/mtp_speculative/dp_proposers/__init__.py | 3 ++- .../server/router/model_infer/mtp_speculative/engine.py | 7 +++++-- .../router/model_infer/mtp_speculative/planner/dspark.py | 7 +++++-- .../model_infer/mtp_speculative/planner/lightspec.py | 7 +++++-- .../model_infer/mtp_speculative/proposers/__init__.py | 3 ++- 11 files changed, 52 insertions(+), 19 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 896e6afee6..82089fc992 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -7,7 +7,7 @@ from sortedcontainers import SortedDict from dataclasses import dataclass, field -from typing import List, Dict, Tuple, Optional, Callable, Any, Union +from typing import TYPE_CHECKING, List, Dict, Tuple, Optional, Callable, Any, Union from lightllm.common.req_manager import ReqManager, ReqManagerForMamba from lightllm.utils.infer_utils import mark_start, mark_end from lightllm.server.core.objs import Req, SamplingParams, FinishStatus, ShmReqManager @@ -25,6 +25,9 @@ from lightllm.server.embed_cache.embed_cache_client import CpuEmbedCacheClient from lightllm.server.router.model_infer.infer_req_ext import FinalTokenMetadataExt, PromptSelectedLogprobsExt +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + logger = init_logger(__name__) @@ -44,15 +47,13 @@ class InferenceContext: def register( self, - backend, + backend: "ModeBackend", req_manager: Union[ReqManager, ReqManagerForMamba], radix_cache: Union[LinearAttPagedRadixCache, RadixCache], shm_req_manager: ShmReqManager, vocab_size: int, ): self.args = get_env_start_args() - from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend - self.backend: ModeBackend = backend self.req_manager = req_manager self.req_sampling_manager = self.req_manager.req_sampling_params_manager diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py index cc5c0e9b43..57692cd55e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py @@ -1,4 +1,6 @@ -from typing import Optional +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional import torch @@ -8,11 +10,14 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + class DPSpecEngine: """普通 DP prefill/decode 使用的 MTP engine。""" - def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: + def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: self.proposer: BaseDpProposer = build_dp_spec_proposer( spec_mode=spec_mode, backend=backend, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index ca7ca8a11f..3fbe894b0b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -1,4 +1,6 @@ -from typing import Optional +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional import torch @@ -11,11 +13,14 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + class DPOverlapSpecEngine: """双 microbatch overlap draft 流程使用的 DP MTP engine。""" - def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: + def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: self.proposer: BaseDpOverlapProposer = build_dp_overlap_spec_proposer( spec_mode=spec_mode, backend=backend, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py index aee2c5b8d3..e85ee8cb34 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py @@ -1,8 +1,13 @@ +from typing import TYPE_CHECKING + from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner.base import BaseDpOverlapPlanner from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner.fixed import FixedDpOverlapPlanner +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + -def build_dp_overlap_planner(*, backend) -> BaseDpOverlapPlanner: +def build_dp_overlap_planner(*, backend: "ModeBackend") -> BaseDpOverlapPlanner: return FixedDpOverlapPlanner(draft_step=backend.max_draft_step) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py index 63841bcaf3..8c0ff84b10 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py @@ -1,13 +1,14 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer def build_dp_overlap_spec_proposer( *, spec_mode: str, - backend, + backend: "ModeBackend", enable_dynmaic_mtp: bool, ) -> "BaseDpOverlapProposer": if spec_mode == "vanilla_with_att": diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py index e29812e24d..d57a0eecd1 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py @@ -1,8 +1,13 @@ +from typing import TYPE_CHECKING + from lightllm.server.router.model_infer.mtp_speculative.dp_planner.base import BaseDpPlanner from lightllm.server.router.model_infer.mtp_speculative.dp_planner.fixed import FixedDpPlanner +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + -def build_dp_planner(*, backend) -> BaseDpPlanner: +def build_dp_planner(*, backend: "ModeBackend") -> BaseDpPlanner: return FixedDpPlanner(draft_step=backend.max_draft_step) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py index 4053a4ee27..1155b67d81 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py @@ -1,10 +1,11 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -def build_dp_spec_proposer(*, spec_mode: str, backend, enable_dynmaic_mtp: bool) -> "BaseDpProposer": +def build_dp_spec_proposer(*, spec_mode: str, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> "BaseDpProposer": if spec_mode == "vanilla_with_att": from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_with_att import ( DpVanillaWithAttProposer, diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 2ed0c3e0bb..6f1dcc0979 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import List, Optional, Tuple +from typing import TYPE_CHECKING, List, Optional, Tuple import torch @@ -16,6 +16,9 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers import build_spec_proposer from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + class SpecEngine: """Owns non-DP MTP planning and draft proposal generation. @@ -24,7 +27,7 @@ class SpecEngine: cleanup are stateless operations exposed by ``mtp_speculative.utils``. """ - def __init__(self, backend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: + def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: self.backend = backend self.proposer: BaseSpecProposer = build_spec_proposer( spec_mode=spec_mode, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index cc1301af97..9f619fbb17 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections import deque -from typing import Dict, List, Optional +from typing import TYPE_CHECKING, Dict, List, Optional import numpy as np @@ -11,6 +11,9 @@ _InferCostMsTable, ) +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + class DSparkPlanner(BaseMtpPlanner): """DSpark's confidence-based verify-capacity planner. @@ -19,7 +22,7 @@ class DSparkPlanner(BaseMtpPlanner): the target verify capacity two iterations later. """ - def __init__(self, backend) -> None: + def __init__(self, backend: ModeBackend) -> None: self.backend = backend self.max_draft_step = int(backend.max_draft_step) self.block_size = int(backend.draft_models[0].block_size) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index c0c70b3172..371a9c792b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Dict, List, Optional, Tuple +from typing import TYPE_CHECKING, Dict, List, Optional, Tuple import numpy as np @@ -11,6 +11,9 @@ _InferCostMsTable, ) +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + class LightSpecPlanner(BaseMtpPlanner): """Choose the current verify budget and the next draft configuration. @@ -37,7 +40,7 @@ class LightSpecPlanner(BaseMtpPlanner): def __init__( self, spec_mode: str, - backend, + backend: ModeBackend, ) -> None: self.spec_mode = spec_mode self.backend = backend diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py index e633095e82..dd0d8d05ac 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py @@ -1,10 +1,11 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer -def build_spec_proposer(*, spec_mode: str, backend, enable_dynmaic_mtp: bool) -> "BaseSpecProposer": +def build_spec_proposer(*, spec_mode: str, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> "BaseSpecProposer": if spec_mode == "dspark": from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer From f3b0b44614931952319c6a22174062f2f3865ab4 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 05:14:30 +0000 Subject: [PATCH 051/103] refactor: centralize MTP memory allocation --- .../dp_overlap_proposers/eagle_utils.py | 5 +++-- .../mtp_speculative/proposers/base.py | 13 ------------- .../mtp_speculative/proposers/eagle_utils.py | 3 ++- .../proposers/parallel_block_utils.py | 3 ++- .../router/model_infer/mtp_speculative/utils.py | 15 +++++++++++++++ .../mtp_speculative/test_eagle_overlap.py | 17 +++++++++++++---- unit_tests/utils/test_speculative_utils.py | 15 ++++++++++++--- 7 files changed, 47 insertions(+), 24 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index bbaa1b5202..d01c18a7b1 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -7,6 +7,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( @@ -165,7 +166,7 @@ def propose_next_dp_eagle_autoregressive_overlap( model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size total_real_request_count = sum(real_request_counts) - extra_mem_indexes_cpu = proposer.alloc_extra_mem_indexes(total_real_request_count * (draft_step - 1)) + extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * (draft_step - 1)) extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) hold_mem_index = proposer.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX @@ -228,7 +229,7 @@ def propose_next_dp_eagle_fixed_layout_overlap( proposal_token_ids[:real_verify_rows0, 0] = next_token_ids0[:real_verify_rows0] proposal_token_ids[real_verify_rows0:, 0] = next_token_ids1[:real_verify_rows1] - extra_mem_indexes_cpu = proposer.alloc_extra_mem_indexes(total_real_request_count * draft_step) + extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * draft_step) extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) split = real_request_counts[0] * draft_step extra_mem_indexes_by_batch = ( diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 0f07062e49..5f783b0164 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -46,19 +46,6 @@ def __init__(self, *, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> None: self.backend = backend self.enable_dynmaic_mtp = bool(enable_dynmaic_mtp) - def alloc_extra_mem_indexes(self, token_count: int) -> torch.Tensor: - """Allocate draft-owned temporary KV slots.""" - - token_count = int(token_count) - if token_count == 0: - return torch.empty((0,), dtype=torch.int32, device="cpu") - - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) - return g_infer_context.req_manager.mem_manager.alloc(token_count) - @abstractmethod def build_draft_state_from_prefill( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 66541e22e2..0af5d269a8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -8,6 +8,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal @@ -135,7 +136,7 @@ def propose_next_eagle( schedule_scores=schedule_scores, ) - extra_mem_indexes_cpu = proposer.alloc_extra_mem_indexes(request_count * (draft_step - 1)) + extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(request_count * (draft_step - 1)) extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) draft_seq_lens = main_model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1 max_kv_seq_len = main_model_input.max_kv_seq_len diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py index de6589a792..d74f0ff28b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py @@ -7,6 +7,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer @@ -58,7 +59,7 @@ def build_parallel_block_draft_input( draft_model = proposer.backend.draft_models[0] block_size = int(draft_model.block_size) - extra_mem_indexes_cpu = proposer.alloc_extra_mem_indexes(request_count * block_size) + extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(request_count * block_size) block_input_ids = next_token_ids.new_full( (request_count * block_size,), diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index 5fbf15a58f..e33f1dc095 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -17,6 +17,20 @@ from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend +def alloc_mem_indexes(token_count: int) -> torch.Tensor: + """Allocate temporary KV slots owned by an MTP proposal.""" + + token_count = int(token_count) + if token_count == 0: + return torch.empty((0,), dtype=torch.int32, device="cpu") + + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + if g_infer_context.radix_cache is not None: + g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(token_count) + return g_infer_context.req_manager.mem_manager.alloc(token_count) + + def verify_mtp_tokens( backend: ModeBackend, next_token_ids: torch.Tensor, @@ -116,6 +130,7 @@ def free_unused_mtp_decode_mem( __all__ = [ + "alloc_mem_indexes", "free_unused_mtp_decode_mem", "record_request_mtp_metrics", "scatter_mtp_next_tokens", diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 0aa2c0378f..ecba9e705b 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -3,6 +3,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import DpOverlapEagle3Proposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_with_att import ( DpOverlapEagleWithAttProposer, @@ -62,7 +63,7 @@ def _target_input(batch_size): ) -def test_overlap_eagle_keeps_fixed_verify_layout(): +def test_overlap_eagle_keeps_fixed_verify_layout(monkeypatch): draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -75,7 +76,11 @@ def test_overlap_eagle_keeps_fixed_verify_layout(): _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=False) - proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) + monkeypatch.setattr( + mtp_utils, + "alloc_mem_indexes", + lambda token_count: torch.arange(token_count, dtype=torch.int32), + ) model_input0 = _target_input(batch_size=6) model_input1 = _target_input(batch_size=6) @@ -109,7 +114,7 @@ def test_overlap_eagle_keeps_fixed_verify_layout(): assert torch.equal(model_input1.mem_indexes, torch.tensor([2, 2, 4, 5, 3, 5], dtype=torch.int32)) -def test_autoregressive_eagle_reuses_overlap_inputs(): +def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -122,7 +127,11 @@ def test_autoregressive_eagle_reuses_overlap_inputs(): _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapEagle3Proposer(backend=backend, enable_dynmaic_mtp=False) - proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) + monkeypatch.setattr( + mtp_utils, + "alloc_mem_indexes", + lambda token_count: torch.arange(token_count, dtype=torch.int32), + ) model_input0 = _target_input(batch_size=6) model_input1 = _target_input(batch_size=6) diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 930390171d..dc2a411b9e 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -13,6 +13,10 @@ from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( + build_parallel_block_draft_input, +) +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.utils import envs_utils @@ -296,7 +300,7 @@ def test_dflash_reuses_decode_input_for_kv_commit(): assert model_input.mtp_draft_input_hiddens is target_hidden -def test_dflash_expands_position_delta_with_request_block_rows(): +def test_dflash_expands_position_delta_with_request_block_rows(monkeypatch): block_size = 3 proposer = DFlashProposer( backend=SimpleNamespace( @@ -304,7 +308,11 @@ def test_dflash_expands_position_delta_with_request_block_rows(): ), enable_dynmaic_mtp=False, ) - proposer.alloc_extra_mem_indexes = lambda token_count: torch.arange(token_count, dtype=torch.int32) + monkeypatch.setattr( + mtp_utils, + "alloc_mem_indexes", + lambda token_count: torch.arange(token_count, dtype=torch.int32), + ) model_input = SimpleNamespace( b_req_idx=torch.tensor([10, 10, 11, 12, 12], dtype=torch.int32), b_seq_len=torch.tensor([4, 5, 7, 8, 9], dtype=torch.int32), @@ -312,7 +320,8 @@ def test_dflash_expands_position_delta_with_request_block_rows(): max_kv_seq_len=9, ) - draft_input, _ = proposer.build_block_draft_input( + draft_input, _ = build_parallel_block_draft_input( + proposer=proposer, main_model_input=model_input, next_token_ids=torch.arange(5, dtype=torch.int64), accepted_tail_rows=torch.tensor([1, 4]), From 6455a15c71781d707a313b0dc72ff9a233ae8283 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 05:39:52 +0000 Subject: [PATCH 052/103] refactor: specialize MTP proposal handling --- .../mode_backend/chunked_prefill/impl.py | 17 +++-- .../mode_backend/dp_backend/impl.py | 12 ++-- .../dp_overlap_proposers/eagle3.py | 6 +- .../dp_overlap_proposers/eagle_no_att.py | 6 +- .../dp_overlap_proposers/eagle_utils.py | 15 +++-- .../dp_overlap_proposers/eagle_with_att.py | 6 +- .../dp_overlap_proposers/vanilla_no_att.py | 6 +- .../dp_overlap_proposers/vanilla_utils.py | 7 +- .../dp_overlap_proposers/vanilla_with_att.py | 6 +- .../mtp_speculative/dp_proposers/eagle3.py | 4 +- .../dp_proposers/eagle_no_att.py | 4 +- .../dp_proposers/eagle_with_att.py | 4 +- .../dp_proposers/vanilla_no_att.py | 4 +- .../dp_proposers/vanilla_with_att.py | 4 +- .../model_infer/mtp_speculative/engine.py | 2 +- .../mtp_speculative/planner/base.py | 11 ++-- .../mtp_speculative/planner/dspark.py | 7 +- .../mtp_speculative/planner/fixed.py | 7 +- .../mtp_speculative/planner/lightspec.py | 3 +- .../mtp_speculative/proposers/base.py | 11 +--- .../mtp_speculative/proposers/dflash.py | 13 +++- .../mtp_speculative/proposers/dspark.py | 14 +++- .../mtp_speculative/proposers/eagle3.py | 5 +- .../mtp_speculative/proposers/eagle_no_att.py | 5 +- .../mtp_speculative/proposers/eagle_utils.py | 16 +++-- .../proposers/eagle_with_att.py | 5 +- .../proposers/vanilla_no_att.py | 5 +- .../proposers/vanilla_utils.py | 13 +++- .../proposers/vanilla_with_att.py | 5 +- .../model_infer/mtp_speculative/utils.py | 26 ++++++-- .../test_dp_overlap_spec_engine.py | 10 +-- .../mtp_speculative/test_planner.py | 66 +++++++++++++++++-- unit_tests/utils/test_speculative_utils.py | 18 +++-- 33 files changed, 235 insertions(+), 108 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 36699d122a..045a5a056c 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -296,6 +296,12 @@ def decode_mtp( verify_event = torch.cuda.Event() verify_event.record() + g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( + b_req_idx=model_input.b_req_idx, + next_token_ids=next_token_ids, + mask=accepted_index == 1, + ) + proposal = spec_engine.propose_next( main_model_input=model_input, main_model_output=model_output, @@ -306,11 +312,10 @@ def decode_mtp( ) mtp_utils.scatter_mtp_next_tokens( backend=self, + proposal=proposal, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=proposal.token_ids, b_req_idx=model_input.b_req_idx, mtp_accept_len=mtp_accept_len, - schedule_scores=proposal.schedule_scores, ) ( @@ -323,12 +328,6 @@ def decode_mtp( next_token_ranks=next_token_ranks, ) - g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( - b_req_idx=model_input.b_req_idx, - next_token_ids=next_token_ids, - mask=accepted_index == 1, - ) - sync_event = torch.cuda.Event() sync_event.record() @@ -377,12 +376,12 @@ def decode_mtp( mtp_utils.free_unused_mtp_decode_mem( backend=self, + proposal=proposal, model_input=model_input, selected_row_mask_cpu=( async_selected_row_mask_cpu.tensor if async_selected_row_mask_cpu is not None else None ), accepted_index_cpu=accepted_index_cpu, - extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, ) # 第四阶段 diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 636f853acd..dabab7cfbb 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -649,10 +649,11 @@ def _draft_decode_vanilla( if req_num > 0: mtp_utils.scatter_mtp_next_tokens( backend=self, - all_next_token_ids=proposal.token_ids[:req_num], + proposal=proposal, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx[:req_num], mtp_accept_len=mtp_accept_len, + valid_row_count=req_num, ) return proposal.extra_mem_indexes_cpu @@ -705,11 +706,11 @@ def _draft_decode_eagle( if req_num > 0: mtp_utils.scatter_mtp_next_tokens( backend=self, + proposal=proposal, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=proposal.token_ids[:req_num], b_req_idx=model_input.b_req_idx[:req_num], mtp_accept_len=mtp_accept_len, - schedule_scores=(proposal.schedule_scores[:req_num] if proposal.schedule_scores is not None else None), + valid_row_count=req_num, ) return proposal.extra_mem_indexes_cpu @@ -988,7 +989,7 @@ def _draft_decode_vanilla_overlap( if req_num0 + req_num1 > 0: mtp_utils.scatter_mtp_next_tokens( backend=self, - all_next_token_ids=proposal.token_ids, + proposal=proposal, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, @@ -1061,10 +1062,9 @@ def _draft_decode_eagle_overlap( if req_num0 + req_num1 > 0: mtp_utils.scatter_mtp_next_tokens( backend=self, + proposal=proposal, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=proposal.token_ids, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, - schedule_scores=proposal.schedule_scores, ) return proposal.extra_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index e3ddd2a3be..b022c4baa8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -6,8 +6,8 @@ build_dp_eagle_draft_state_from_prefill_overlap, propose_next_dp_eagle_autoregressive_overlap, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, propose_next_eagle, ) @@ -54,7 +54,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( self, main_model_input, @@ -79,7 +79,7 @@ def propose_next_overlap( real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_dp_eagle_autoregressive_overlap( self, main_model_input0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 7e147a0106..5029b27e8c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -6,8 +6,8 @@ build_dp_eagle_draft_state_from_prefill_overlap, propose_next_dp_eagle_fixed_layout_overlap, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, propose_next_eagle, ) @@ -51,7 +51,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( self, main_model_input, @@ -76,7 +76,7 @@ def propose_next_overlap( real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_dp_eagle_fixed_layout_overlap( self, main_model_input0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index d01c18a7b1..411908fa35 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -9,8 +9,8 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, generate_eagle_token_ids, prepare_eagle_verify_extend_input, ) @@ -71,7 +71,7 @@ def propose_next_dp_eagle_autoregressive_overlap( accept_len1: torch.Tensor | None, draft_step: int, map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> SpecProposal: +) -> EagleSpecProposal: """运行 DP EAGLE extend 后接单 token overlap decode 的 proposal 流程。""" verify_width = proposer.backend.max_draft_step + 1 @@ -143,9 +143,10 @@ def propose_next_dp_eagle_autoregressive_overlap( proposal_token_ids[proposal_rows, 1] = draft_token_ids[:real_request_count] if draft_step == 1: - return SpecProposal( + return EagleSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, + schedule_scores=None, ) for batch_index, model_input in enumerate(model_inputs): @@ -197,9 +198,10 @@ def propose_next_dp_eagle_autoregressive_overlap( real_request_count = real_request_counts[batch_index] proposal_token_ids[proposal_rows_by_batch[batch_index], step + 1] = draft_token_ids[:real_request_count] - return SpecProposal( + return EagleSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=extra_mem_indexes_cpu, + schedule_scores=None, ) @@ -215,7 +217,7 @@ def propose_next_dp_eagle_fixed_layout_overlap( real_verify_rows1: int, draft_step: int, map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> SpecProposal: +) -> EagleSpecProposal: """保持 expanded verify-row layout 运行 DP EAGLE overlap decode。""" verify_width = proposer.backend.max_draft_step + 1 @@ -282,7 +284,8 @@ def propose_next_dp_eagle_fixed_layout_overlap( proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids_by_batch[0][:real_verify_rows0] proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids_by_batch[1][:real_verify_rows1] - return SpecProposal( + return EagleSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=extra_mem_indexes_cpu, + schedule_scores=None, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index d3e51c78f5..8e15e1ee1d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -6,8 +6,8 @@ build_dp_eagle_draft_state_from_prefill_overlap, propose_next_dp_eagle_fixed_layout_overlap, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, propose_next_eagle, ) @@ -51,7 +51,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( self, main_model_input, @@ -76,7 +76,7 @@ def propose_next_overlap( real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_dp_eagle_fixed_layout_overlap( self, main_model_input0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 6fa226a5e4..7ddf7e22ee 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -6,8 +6,8 @@ build_dp_chained_mtp_draft_state_from_prefill_overlap, propose_next_dp_chained_mtp_overlap, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + VanillaSpecProposal, build_chained_mtp_draft_state_from_prefill, propose_next_chained_mtp, ) @@ -51,7 +51,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> VanillaSpecProposal: return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) def propose_next_overlap( @@ -67,7 +67,7 @@ def propose_next_overlap( real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, - ) -> SpecProposal: + ) -> VanillaSpecProposal: return propose_next_dp_chained_mtp_overlap( self, main_model_input0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py index b60066fd1b..ce4d9307d2 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py @@ -6,7 +6,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal def build_dp_chained_mtp_draft_state_from_prefill_overlap( @@ -53,7 +53,7 @@ def propose_next_dp_chained_mtp_overlap( next_token_ids1: torch.Tensor, real_verify_rows1: int, draft_step: int, -) -> SpecProposal: +) -> VanillaSpecProposal: """为两个 DP microbatch 运行 Vanilla chained overlap decode。""" model_inputs = (main_model_input0, main_model_input1) @@ -80,7 +80,8 @@ def propose_next_dp_chained_mtp_overlap( proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids[0][:real_verify_rows0] proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids[1][:real_verify_rows1] - return SpecProposal( + return VanillaSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, + schedule_scores=None, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 642dc8848f..6050d544fa 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -6,8 +6,8 @@ build_dp_chained_mtp_draft_state_from_prefill_overlap, propose_next_dp_chained_mtp_overlap, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + VanillaSpecProposal, build_chained_mtp_draft_state_from_prefill, propose_next_chained_mtp, ) @@ -51,7 +51,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> VanillaSpecProposal: return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) def propose_next_overlap( @@ -67,7 +67,7 @@ def propose_next_overlap( real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, - ) -> SpecProposal: + ) -> VanillaSpecProposal: return propose_next_dp_chained_mtp_overlap( self, main_model_input0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py index 6852354716..d4289e0060 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py @@ -2,8 +2,8 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, propose_next_eagle, ) @@ -31,7 +31,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( self, main_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py index d809363cb4..637eef849d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py @@ -2,8 +2,8 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, propose_next_eagle, ) @@ -28,7 +28,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( self, main_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py index b3cbbc9f12..0bdfa38be6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py @@ -2,8 +2,8 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, propose_next_eagle, ) @@ -28,7 +28,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( self, main_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py index 83f42a07c2..2ec72367d0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -2,8 +2,8 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + VanillaSpecProposal, build_chained_mtp_draft_state_from_prefill, propose_next_chained_mtp, ) @@ -28,5 +28,5 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> VanillaSpecProposal: return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py index c08e09b2fd..9b1bf82be5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -2,8 +2,8 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + VanillaSpecProposal, build_chained_mtp_draft_state_from_prefill, propose_next_chained_mtp, ) @@ -28,5 +28,5 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> VanillaSpecProposal: return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 6f1dcc0979..107fed3e1e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -126,9 +126,9 @@ def update_planner_statics( self.planner.update_statics( plan=plan, + proposal=proposal, req_num=req_num, accept_lengths=accept_lengths_cpu, - schedule_scores=proposal.schedule_scores_cpu, ) def _build_mtp_planner(self, spec_mode: str, enable_dynmaic_mtp: bool) -> BaseMtpPlanner: diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py index e278157cd0..27f790c54f 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -2,10 +2,13 @@ from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import List +from typing import TYPE_CHECKING, List from sortedcontainers import SortedDict +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + @dataclass(frozen=True) class SpecDecodePlan: @@ -62,19 +65,19 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: def update_statics( self, plan: SpecDecodePlan, + proposal: SpecProposal, req_num: int, accept_lengths, - schedule_scores=None, ) -> None: """在本轮 verify 完成后更新规划器的运行时统计。 Args: plan: 本轮 decode 实际采用的执行计划,用于确定被验证的配置。 + proposal: proposer 生成的模式专属输出。需要额外调度信息的规划器 + 直接读取自己的 proposal 子类,其他规划器忽略该对象。 req_num: 本轮逻辑请求数量。 accept_lengths: 每个请求本轮提交的 token 数量,包含必然提交的 target token。 - schedule_scores: proposer 可选提供的调度分数。LightSpec 主要使用 - accept_lengths,DSpark 使用置信度分数,固定规划器忽略所有反馈。 """ raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index 9f619fbb17..a04134c2e3 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -10,6 +10,7 @@ SpecDecodePlan, _InferCostMsTable, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkSpecProposal if TYPE_CHECKING: from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend @@ -52,13 +53,13 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: def update_statics( self, plan: SpecDecodePlan, + proposal: DSparkSpecProposal, req_num: int, accept_lengths, - schedule_scores=None, ) -> None: - if schedule_scores is not None: + if proposal.schedule_scores_cpu is not None: self._update_confidence_probs( - confidence_probs=schedule_scores, + confidence_probs=proposal.schedule_scores_cpu, req_num=req_num, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py index 34976164db..b46fdd1ceb 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py @@ -1,9 +1,12 @@ from __future__ import annotations -from typing import List +from typing import TYPE_CHECKING, List from lightllm.server.router.model_infer.mtp_speculative.planner.base import BaseMtpPlanner, SpecDecodePlan +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + class FixedSpecPlanner(BaseMtpPlanner): """Planner for fixed-width speculative decoding.""" @@ -22,9 +25,9 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: def update_statics( self, plan: SpecDecodePlan, + proposal: SpecProposal, req_num: int, accept_lengths, - schedule_scores=None, ) -> None: """固定规划不根据运行反馈调整 batch size 或 draft step。""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index 371a9c792b..b190285188 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -13,6 +13,7 @@ if TYPE_CHECKING: from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal class LightSpecPlanner(BaseMtpPlanner): @@ -114,9 +115,9 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: def update_statics( self, plan: SpecDecodePlan, + proposal: SpecProposal, req_num: int, accept_lengths, - schedule_scores=None, ) -> None: # The progress EMA records one complete-batch sample for a single # (N, B, d) configuration. all_reqs_have_proposals is false if any request is diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 5f783b0164..f3258a646c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -14,23 +14,16 @@ @dataclass class SpecProposal: - """Candidate tokens and scheduling metadata produced by a proposer. + """Common candidate-token output produced by every MTP proposer. `token_ids` has shape `[verify_batch, draft_step + 1]`; column 0 contains target-model tokens and the remaining columns contain draft candidates. - `schedule_scores`, when present, has shape `[verify_batch, draft_step]`. - Each column contains the proposer-specific score used by dynamic scheduling: - selected-token probability for standard proposers, or confidence-head - probability for DSpark. - `schedule_scores_cpu` is the asynchronous CPU copy consumed by planners - that use proposal scores directly. `extra_mem_indexes_cpu` tracks temporary KV slots owned by the proposal. + Mode-specific scheduling metadata belongs to the corresponding subclass. """ token_ids: torch.Tensor extra_mem_indexes_cpu: Optional[torch.Tensor] - schedule_scores: Optional[torch.Tensor] = None - schedule_scores_cpu: Optional[torch.Tensor] = None class BaseSpecProposer(ABC): diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index 1311bcaa26..26fb515acc 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -1,5 +1,7 @@ from __future__ import annotations +from dataclasses import dataclass + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput @@ -11,6 +13,13 @@ ) +@dataclass +class DFlashSpecProposal(SpecProposal): + """DFlash proposal with optional block-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + class DFlashProposer(BaseSpecProposer): """DFlash block-diffusion proposer. @@ -40,7 +49,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> DFlashSpecProposal: request_count = int(b_req_mtp_start_loc.shape[0]) verify_row_count = int(next_token_ids.shape[0]) draft_model = self.backend.draft_models[0] @@ -83,7 +92,7 @@ def propose_next( device=next_token_ids.device, ) schedule_scores[accepted_tail_rows] = block_draft_token_probs[:, :draft_step].float() - return SpecProposal( + return DFlashSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=extra_mem_indexes_cpu, schedule_scores=schedule_scores, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 71f59c6ddb..48cfd2a0fc 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -1,5 +1,7 @@ from __future__ import annotations +from dataclasses import dataclass + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput @@ -12,6 +14,14 @@ ) +@dataclass +class DSparkSpecProposal(SpecProposal): + """DSpark proposal with GPU confidence scores and their CPU planner view.""" + + schedule_scores: torch.Tensor | None = None + schedule_scores_cpu: torch.Tensor | None = None + + class DSparkProposer(BaseSpecProposer): """DSpark semi-autoregressive parallel-block proposer. @@ -42,7 +52,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> DSparkSpecProposal: request_count = int(b_req_mtp_start_loc.shape[0]) verify_row_count = int(next_token_ids.shape[0]) draft_model = self.backend.draft_models[0] @@ -106,7 +116,7 @@ def propose_next( gpu_tensor=schedule_scores, ) - return SpecProposal( + return DSparkSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=extra_mem_indexes_cpu, schedule_scores=schedule_scores, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py index 8ffe8cd477..2c618bc98d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py @@ -1,8 +1,9 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, generate_eagle_token_ids, generate_eagle_token_ids_and_prob, @@ -38,7 +39,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( proposer=self, main_model_input=main_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py index d74e8b87ee..bf1896e83e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -1,8 +1,9 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, propose_next_eagle, ) @@ -27,7 +28,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( proposer=self, main_model_input=main_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 0af5d269a8..b0fb83acde 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -3,6 +3,7 @@ from __future__ import annotations import copy +from dataclasses import dataclass from typing import Callable import torch @@ -12,6 +13,13 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +@dataclass +class EagleSpecProposal(SpecProposal): + """EAGLE proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + def build_eagle_draft_state_from_prefill( proposer: BaseSpecProposer, target_model_input: ModelInput, @@ -75,7 +83,7 @@ def propose_next_eagle( draft_step: int, accept_len: torch.Tensor | None, map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> SpecProposal: +) -> EagleSpecProposal: """运行 EAGLE extend 后接单 token decode 的通用 proposal 流程。""" verify_row_count = int(next_token_ids.shape[0]) @@ -96,7 +104,7 @@ def propose_next_eagle( else None ) if draft_step == 0: - return SpecProposal( + return EagleSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, schedule_scores=schedule_scores, @@ -130,7 +138,7 @@ def propose_next_eagle( draft_hidden = extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows) if draft_step == 1: - return SpecProposal( + return EagleSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, schedule_scores=schedule_scores, @@ -180,7 +188,7 @@ def propose_next_eagle( draft_hidden = draft_output.mtp_collector.spec_hidden draft_seq_lens.add_(1) - return SpecProposal( + return EagleSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=extra_mem_indexes_cpu, schedule_scores=schedule_scores, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index b66f809925..ab0cb63f5a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -1,8 +1,9 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( + EagleSpecProposal, build_eagle_draft_state_from_prefill, propose_next_eagle, ) @@ -27,7 +28,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> EagleSpecProposal: return propose_next_eagle( proposer=self, main_model_input=main_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index 12f5f48c61..13a1efa67e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -1,8 +1,9 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + VanillaSpecProposal, build_chained_mtp_draft_state_from_prefill, propose_next_chained_mtp, ) @@ -32,7 +33,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> VanillaSpecProposal: return propose_next_chained_mtp( proposer=self, main_model_input=main_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py index 23718029f9..573931d5c6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py @@ -2,12 +2,21 @@ from __future__ import annotations +from dataclasses import dataclass + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +@dataclass +class VanillaSpecProposal(SpecProposal): + """Vanilla proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + def build_chained_mtp_draft_state_from_prefill( proposer: BaseSpecProposer, target_model_input: ModelInput, @@ -37,7 +46,7 @@ def propose_next_chained_mtp( main_model_output: ModelOutput, next_token_ids: torch.Tensor, draft_step: int, -) -> SpecProposal: +) -> VanillaSpecProposal: """依次运行 Vanilla chained MTP 模块并生成 proposal。""" verify_row_count = int(next_token_ids.shape[0]) @@ -68,7 +77,7 @@ def propose_next_chained_mtp( draft_token_ids = proposer.backend._gen_argmax_token_ids(draft_output) proposal_token_ids[:, step + 1] = draft_token_ids - return SpecProposal( + return VanillaSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=None, schedule_scores=schedule_scores, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index 5807c5dfdf..23fc7eaeb1 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -1,8 +1,9 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( + VanillaSpecProposal, build_chained_mtp_draft_state_from_prefill, propose_next_chained_mtp, ) @@ -32,7 +33,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - ) -> SpecProposal: + ) -> VanillaSpecProposal: return propose_next_chained_mtp( proposer=self, main_model_input=main_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index e33f1dc095..7681e26858 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -15,6 +15,7 @@ if TYPE_CHECKING: from lightllm.server.router.model_infer.infer_batch import InferReq from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal def alloc_mem_indexes(token_count: int) -> torch.Tensor: @@ -61,19 +62,32 @@ def verify_mtp_tokens( def scatter_mtp_next_tokens( backend: ModeBackend, + proposal: SpecProposal, b_req_mtp_start_loc: torch.Tensor, - all_next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, mtp_accept_len: torch.Tensor, - schedule_scores: Optional[torch.Tensor] = None, + valid_row_count: Optional[int] = None, ) -> None: """Persist the next MTP proposal and optional scheduling scores by request.""" + from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashSpecProposal + from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkSpecProposal + from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import EagleSpecProposal + from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal + + token_ids = proposal.token_ids + scored_proposal_types = (VanillaSpecProposal, EagleSpecProposal, DFlashSpecProposal, DSparkSpecProposal) + schedule_scores = proposal.schedule_scores if isinstance(proposal, scored_proposal_types) else None + if valid_row_count is not None: + token_ids = token_ids[:valid_row_count] + if schedule_scores is not None: + schedule_scores = schedule_scores[:valid_row_count] + sampling_params_manager = backend.model.req_manager.req_sampling_params_manager mtp_scatter_next_token_ids( req_to_next_token_ids=sampling_params_manager.req_to_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, + all_next_token_ids=token_ids, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, req_to_next_token_scores=( @@ -107,10 +121,10 @@ def record_request_mtp_metrics( def free_unused_mtp_decode_mem( backend: ModeBackend, + proposal: SpecProposal, model_input: ModelInput, selected_row_mask_cpu: Optional[torch.Tensor], accepted_index_cpu: torch.Tensor, - extra_mem_indexes_cpu: Optional[torch.Tensor], ) -> None: """Free rejected target KV slots and draft-only temporary slots.""" @@ -123,8 +137,8 @@ def free_unused_mtp_decode_mem( free_mask[selected_mask] = accepted_index_cpu == 0 need_free_mem_indexes = mem_indexes_cpu[free_mask] - if extra_mem_indexes_cpu is not None: - need_free_mem_indexes = torch.cat([need_free_mem_indexes, extra_mem_indexes_cpu], dim=0) + if proposal.extra_mem_indexes_cpu is not None: + need_free_mem_indexes = torch.cat([need_free_mem_indexes, proposal.extra_mem_indexes_cpu], dim=0) if len(need_free_mem_indexes) > 0: backend.model.req_manager.mem_manager.free(need_free_mem_indexes) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 1449ed7162..8681e37f58 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -124,7 +124,8 @@ def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(monkeypatch): assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) assert torch.equal(propose_args["b_req_mtp_start_loc"], torch.tensor([0, 8], dtype=torch.int32)) assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) - assert scatter_args["all_next_token_ids"].shape == (8, 8) + assert scatter_args["proposal"].token_ids.shape == (16, 8) + assert scatter_args["valid_row_count"] == 8 assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) @@ -152,7 +153,8 @@ def test_dp_vanilla_uses_dp_engine_proposer(monkeypatch): propose_args = backend.spec_engine.propose_args assert propose_args["next_token_ids"].shape == (16,) assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) - assert scatter_args["all_next_token_ids"].shape == (8, 8) + assert scatter_args["proposal"].token_ids.shape == (16, 8) + assert scatter_args["valid_row_count"] == 8 assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) @@ -198,7 +200,7 @@ def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(monkeypat assert propose_args["real_verify_rows1"] == 16 assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) - assert scatter_args["all_next_token_ids"].shape == (24, 8) + assert scatter_args["proposal"].token_ids.shape == (24, 8) assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) @@ -239,5 +241,5 @@ def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(monkeypatch) assert propose_args["real_verify_rows1"] == 16 assert propose_args["accept_len0"] is None assert propose_args["accept_len1"] is None - assert scatter_args["all_next_token_ids"].shape == (24, 8) + assert scatter_args["proposal"].token_ids.shape == (24, 8) assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 8a4e814d48..7ce4ca3e3b 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -51,12 +51,14 @@ from lightllm.server.router.model_infer.mtp_speculative.planner.base import _InferCostMsTable from lightllm.server.router.model_infer.mtp_speculative.proposers import build_spec_proposer from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal -from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer, DFlashSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer, DSparkSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import EagleNoAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import EagleSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import VanillaWithAttProposer from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils @@ -139,6 +141,58 @@ def test_spec_engine_only_exposes_planning_and_proposal_interfaces(): } +def test_mode_proposals_own_their_schedule_metadata(): + assert "schedule_scores" not in SpecProposal.__dataclass_fields__ + assert "schedule_scores_cpu" not in SpecProposal.__dataclass_fields__ + + for proposal_type in (VanillaSpecProposal, EagleSpecProposal, DFlashSpecProposal, DSparkSpecProposal): + assert "schedule_scores" in proposal_type.__dataclass_fields__ + assert "schedule_scores_cpu" not in VanillaSpecProposal.__dataclass_fields__ + assert "schedule_scores_cpu" not in EagleSpecProposal.__dataclass_fields__ + assert "schedule_scores_cpu" not in DFlashSpecProposal.__dataclass_fields__ + assert "schedule_scores_cpu" in DSparkSpecProposal.__dataclass_fields__ + + +def test_scatter_mtp_next_tokens_consumes_mode_proposal(monkeypatch): + scatter_args = {} + monkeypatch.setattr( + mtp_utils, + "mtp_scatter_next_token_ids", + lambda **kwargs: scatter_args.update(kwargs), + ) + req_to_next_token_ids = torch.empty((4, 2), dtype=torch.int64) + req_to_next_token_scores = torch.empty((4, 1), dtype=torch.float32) + backend = SimpleNamespace( + model=SimpleNamespace( + req_manager=SimpleNamespace( + req_sampling_params_manager=SimpleNamespace( + req_to_next_token_ids=req_to_next_token_ids, + req_to_next_token_scores=req_to_next_token_scores, + ) + ) + ) + ) + proposal = DFlashSpecProposal( + token_ids=torch.arange(6, dtype=torch.int64).view(3, 2), + extra_mem_indexes_cpu=None, + schedule_scores=torch.arange(3, dtype=torch.float32).view(3, 1), + ) + + mtp_utils.scatter_mtp_next_tokens( + backend=backend, + proposal=proposal, + b_req_mtp_start_loc=torch.tensor([0, 1], dtype=torch.int32), + b_req_idx=torch.tensor([0, 1], dtype=torch.int32), + mtp_accept_len=torch.ones(2, dtype=torch.int32), + valid_row_count=2, + ) + + assert scatter_args["req_to_next_token_ids"] is req_to_next_token_ids + assert scatter_args["req_to_next_token_scores"] is req_to_next_token_scores + assert torch.equal(scatter_args["all_next_token_ids"], proposal.token_ids[:2]) + assert torch.equal(scatter_args["schedule_scores"], proposal.schedule_scores[:2]) + + def test_dp_planner_returns_fixed_backend_draft_step(): planner = build_dp_planner(backend=SimpleNamespace(max_draft_step=4)) @@ -306,6 +360,7 @@ def test_eagle_proposer_skips_draft_forward_for_zero_steps(): draft_step=0, ) + assert isinstance(proposal, EagleSpecProposal) assert proposal.token_ids.tolist() == [[10], [11]] assert proposal.schedule_scores.shape == (2, 0) @@ -321,10 +376,13 @@ def test_dynamic_decode_frees_unselected_and_rejected_rows(): mtp_utils.free_unused_mtp_decode_mem( backend=backend, + proposal=SpecProposal( + token_ids=torch.empty((0,), dtype=torch.int64), + extra_mem_indexes_cpu=torch.tensor([20]), + ), model_input=model_input, selected_row_mask_cpu=torch.tensor([1, 0, 1, 0], dtype=torch.bool), accepted_index_cpu=torch.tensor([1, 0], dtype=torch.int32), - extra_mem_indexes_cpu=torch.tensor([20]), ) assert len(freed) == 1 @@ -695,7 +753,7 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): planner.draft_infer_costs.update(batch_size=6, infer_cost_ms=0.5) confidence_probs = np.asarray([[0.9, 0.9, 0.9]] * 2, dtype=np.float64) plan = SpecDecodePlan(origin_batch_size=8, dynamic_batch_size=8, draft_step=3, pre_draft_step=3) - proposal = SpecProposal( + proposal = DSparkSpecProposal( token_ids=torch.empty((0,), dtype=torch.int64), extra_mem_indexes_cpu=None, schedule_scores_cpu=torch.from_numpy(confidence_probs), diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index dc2a411b9e..0b57b2e86d 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -12,7 +12,8 @@ from lightllm.common.basemodel.attention.fa3.mla import MlaFa3DecodeAttState, MlaFa3PrefillAttState from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight -from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers import dflash as dflash_module +from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer, DFlashSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( build_parallel_block_draft_input, ) @@ -223,7 +224,7 @@ def test_draft_model_registry_rejects_unsupported_mode(model_type, spec_mode): ) -def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): +def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(monkeypatch): block_size = 4 max_draft_step = 3 verify_row_count = 5 @@ -241,18 +242,25 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(): _gen_argmax_token_ids_and_prob=lambda _: (flat_draft_token_ids, flat_draft_token_probs), ) proposer = DFlashProposer(backend=backend, enable_dynmaic_mtp=True) - proposer.extend_draft_kv_cache = lambda **_: None - proposer.build_block_draft_input = lambda **_: (SimpleNamespace(), torch.tensor([10, 11])) + monkeypatch.setattr( + dflash_module, + "build_parallel_block_draft_input", + lambda **_: (SimpleNamespace(), torch.tensor([10, 11])), + ) + monkeypatch.setattr(dflash_module, "extend_parallel_block_draft_kv_cache", lambda **_: None) proposal = proposer.propose_next( main_model_input=SimpleNamespace(), - main_model_output=SimpleNamespace(spec_hidden=torch.empty(verify_row_count, 1)), + main_model_output=SimpleNamespace( + mtp_collector=SimpleNamespace(spec_hidden=torch.empty(verify_row_count, 1)), + ), next_token_ids=torch.arange(verify_row_count), b_req_mtp_start_loc=torch.tensor([0, 3]), draft_step=2, accept_len=torch.tensor([1, 1]), ) + assert isinstance(proposal, DFlashSpecProposal) expected_blocks = flat_draft_token_ids.reshape(2, block_size)[:, :2] torch.testing.assert_close(proposal.token_ids[accepted_tail_rows, 1:], expected_blocks) assert proposal.schedule_scores.shape == (verify_row_count, 2) From accdcdf1688d7f6374c99855f27fc50a02dae5eb Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 05:57:49 +0000 Subject: [PATCH 053/103] refactor: unify MTP proposal memory release --- .../mode_backend/dp_backend/impl.py | 22 ++++++----- .../dp_overlap_proposers/eagle_utils.py | 7 ++-- .../dp_overlap_proposers/vanilla_utils.py | 2 +- .../mtp_speculative/proposers/base.py | 27 ++++++++++++-- .../mtp_speculative/proposers/dflash.py | 8 +++- .../mtp_speculative/proposers/dspark.py | 8 +++- .../mtp_speculative/proposers/eagle_utils.py | 12 ++++-- .../proposers/vanilla_utils.py | 2 +- .../model_infer/mtp_speculative/utils.py | 37 ++++++++++++++++--- .../test_dp_overlap_spec_engine.py | 23 ++++++++---- .../mtp_speculative/test_eagle_overlap.py | 8 +++- .../mtp_speculative/test_planner.py | 22 ++++++++--- .../mtp_speculative/test_vanilla_overlap.py | 2 +- 13 files changed, 132 insertions(+), 48 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index dabab7cfbb..8baf032a7a 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -566,7 +566,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): verify_event = torch.cuda.Event() verify_event.record() - eagle_mem_indexes_cpu = self._draft_decode_func( + extra_mem_indexes_cpu = self._draft_decode_func( model_input=model_input, model_output=model_output, next_token_ids=next_token_ids, @@ -601,8 +601,6 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() need_free_mem_indexes = model_input.mem_indexes_cpu[0:req_num][accepted_index_cpu == 0] - if eagle_mem_indexes_cpu is not None: - need_free_mem_indexes = torch.cat([need_free_mem_indexes, eagle_mem_indexes_cpu], dim=0) select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( @@ -613,8 +611,11 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): run_reqs_update_packs=update_packs, extra_post_req_handle_func=self.extra_post_req_handle_func, ) - if len(need_free_mem_indexes) > 0: - g_infer_context.req_manager.mem_manager.free(need_free_mem_indexes) + mtp_utils.free_mem_indexes( + backend=self, + mem_indexes_cpu=need_free_mem_indexes, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, + ) # 第四阶段 event_pack.notify_pre_post_handle() @@ -884,7 +885,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf verify_event = torch.cuda.Event() verify_event.record() - eagle_mem_indexes_cpu = self._draft_decode_overlap_func( + extra_mem_indexes_cpu = self._draft_decode_overlap_func( model_input0=model_input0, model_input1=model_input1, model_output0=model_output0, @@ -924,8 +925,6 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf (model_input0.mem_indexes_cpu[0:req_num0], model_input1.mem_indexes_cpu[0:req_num1]), dim=0 ) need_free_mem_indexes = mem_indexes_cpu[accepted_index_cpu == 0] - if eagle_mem_indexes_cpu is not None: - need_free_mem_indexes = torch.cat((need_free_mem_indexes, eagle_mem_indexes_cpu), dim=0) select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( @@ -936,8 +935,11 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf run_reqs_update_packs=update_packs, extra_post_req_handle_func=self.extra_post_req_handle_func, ) - if len(need_free_mem_indexes) > 0: - g_infer_context.req_manager.mem_manager.free(need_free_mem_indexes) + mtp_utils.free_mem_indexes( + backend=self, + mem_indexes_cpu=need_free_mem_indexes, + extra_mem_indexes_cpu=extra_mem_indexes_cpu, + ) event_pack.notify_pre_post_handle() else: event_pack.notify_post_handle_and_wait_pre_post_handle() diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 411908fa35..3724edf8a9 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -9,6 +9,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, generate_eagle_token_ids, @@ -145,7 +146,7 @@ def propose_next_dp_eagle_autoregressive_overlap( if draft_step == 1: return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, + extra_mem_indexes_cpu=[], schedule_scores=None, ) @@ -200,7 +201,7 @@ def propose_next_dp_eagle_autoregressive_overlap( return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], schedule_scores=None, ) @@ -286,6 +287,6 @@ def propose_next_dp_eagle_fixed_layout_overlap( return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], schedule_scores=None, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py index ce4d9307d2..5a663dab29 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py @@ -82,6 +82,6 @@ def propose_next_dp_chained_mtp_overlap( return VanillaSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, + extra_mem_indexes_cpu=[], schedule_scores=None, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index f3258a646c..b8e16fc030 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -1,8 +1,8 @@ from __future__ import annotations from abc import ABC, abstractmethod -from dataclasses import dataclass -from typing import TYPE_CHECKING, Optional +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, List, Optional import torch @@ -12,18 +12,37 @@ from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend +@dataclass +class MtpMemIndexesToFree: + """描述一组由 MTP proposal 持有、需要在 verify 后释放的临时 KV 索引。""" + + # 位于 CPU 上的临时 KV cache 索引。 + mem_indexes_cpu: torch.Tensor + # 与 mem_indexes_cpu 形状一致的 bool Tensor;True 表示释放对应索引。 + # 为 None 时表示 mem_indexes_cpu 中的全部索引都需要释放。 + free_mask_cpu: Optional[torch.Tensor] = None + + def __post_init__(self) -> None: + if self.free_mask_cpu is None: + return + assert isinstance(self.free_mask_cpu, torch.Tensor) + assert self.free_mask_cpu.dtype == torch.bool + assert self.free_mask_cpu.shape == self.mem_indexes_cpu.shape + + @dataclass class SpecProposal: """Common candidate-token output produced by every MTP proposer. `token_ids` has shape `[verify_batch, draft_step + 1]`; column 0 contains target-model tokens and the remaining columns contain draft candidates. - `extra_mem_indexes_cpu` tracks temporary KV slots owned by the proposal. + `extra_mem_indexes_cpu` tracks temporary KV slots owned by the proposal; + each item describes one independently managed group of temporary indexes. Mode-specific scheduling metadata belongs to the corresponding subclass. """ token_ids: torch.Tensor - extra_mem_indexes_cpu: Optional[torch.Tensor] + extra_mem_indexes_cpu: List[MtpMemIndexesToFree] = field(default_factory=list) class BaseSpecProposer(ABC): diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index 26fb515acc..ecfb218080 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -5,7 +5,11 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + BaseSpecProposer, + MtpMemIndexesToFree, + SpecProposal, +) from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( build_parallel_block_draft_input, build_parallel_block_draft_state_from_prefill, @@ -94,6 +98,6 @@ def propose_next( schedule_scores[accepted_tail_rows] = block_draft_token_probs[:, :draft_step].float() return DFlashSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 48cfd2a0fc..02134ede99 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -6,7 +6,11 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + BaseSpecProposer, + MtpMemIndexesToFree, + SpecProposal, +) from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( build_parallel_block_draft_input, build_parallel_block_draft_state_from_prefill, @@ -118,7 +122,7 @@ def propose_next( return DSparkSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], schedule_scores=schedule_scores, schedule_scores_cpu=schedule_scores_cpu, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index b0fb83acde..561e544bea 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -10,7 +10,11 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + BaseSpecProposer, + MtpMemIndexesToFree, + SpecProposal, +) @dataclass @@ -106,7 +110,7 @@ def propose_next_eagle( if draft_step == 0: return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, + extra_mem_indexes_cpu=[], schedule_scores=schedule_scores, ) @@ -140,7 +144,7 @@ def propose_next_eagle( if draft_step == 1: return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, + extra_mem_indexes_cpu=[], schedule_scores=schedule_scores, ) @@ -190,6 +194,6 @@ def propose_next_eagle( return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py index 573931d5c6..0d4953d9cd 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py @@ -79,6 +79,6 @@ def propose_next_chained_mtp( return VanillaSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=None, + extra_mem_indexes_cpu=[], schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index 7681e26858..01789171d0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -15,7 +15,10 @@ if TYPE_CHECKING: from lightllm.server.router.model_infer.infer_batch import InferReq from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend - from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + MtpMemIndexesToFree, + SpecProposal, + ) def alloc_mem_indexes(token_count: int) -> torch.Tensor: @@ -119,6 +122,28 @@ def record_request_mtp_metrics( req.update_mtp_verify_step_num(verify_step_num=1) +def free_mem_indexes( + backend: ModeBackend, + mem_indexes_cpu: torch.Tensor, + extra_mem_indexes_cpu: List[MtpMemIndexesToFree], +) -> None: + """Free regular KV indexes plus proposal-owned indexes selected by masks.""" + + mem_indexes_to_free = [] + if mem_indexes_cpu.numel() > 0: + mem_indexes_to_free.append(mem_indexes_cpu) + + for extra_mem_to_free in extra_mem_indexes_cpu: + extra_indexes_cpu = extra_mem_to_free.mem_indexes_cpu + if extra_mem_to_free.free_mask_cpu is not None: + extra_indexes_cpu = extra_indexes_cpu[extra_mem_to_free.free_mask_cpu] + if extra_indexes_cpu.numel() > 0: + mem_indexes_to_free.append(extra_indexes_cpu) + + if mem_indexes_to_free: + backend.model.req_manager.mem_manager.free(torch.cat(mem_indexes_to_free, dim=0)) + + def free_unused_mtp_decode_mem( backend: ModeBackend, proposal: SpecProposal, @@ -137,14 +162,16 @@ def free_unused_mtp_decode_mem( free_mask[selected_mask] = accepted_index_cpu == 0 need_free_mem_indexes = mem_indexes_cpu[free_mask] - if proposal.extra_mem_indexes_cpu is not None: - need_free_mem_indexes = torch.cat([need_free_mem_indexes, proposal.extra_mem_indexes_cpu], dim=0) - if len(need_free_mem_indexes) > 0: - backend.model.req_manager.mem_manager.free(need_free_mem_indexes) + free_mem_indexes( + backend=backend, + mem_indexes_cpu=need_free_mem_indexes, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, + ) __all__ = [ "alloc_mem_indexes", + "free_mem_indexes", "free_unused_mtp_decode_mem", "record_request_mtp_metrics", "scatter_mtp_next_tokens", diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 8681e37f58..b7c724fbf2 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -9,7 +9,10 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + MtpMemIndexesToFree, + SpecProposal, +) class _RecordingSpecEngine: @@ -22,7 +25,7 @@ def propose_next(self, **kwargs): token_ids = kwargs["next_token_ids"].new_zeros((16, 8)) return SpecProposal( token_ids=token_ids, - extra_mem_indexes_cpu=torch.tensor([123], dtype=torch.int32), + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([123], dtype=torch.int32))], ) def propose_next_overlap(self, **kwargs): @@ -31,7 +34,7 @@ def propose_next_overlap(self, **kwargs): token_ids = kwargs["next_token_ids0"].new_zeros((row_count, 8)) return SpecProposal( token_ids=token_ids, - extra_mem_indexes_cpu=torch.tensor([456], dtype=torch.int32), + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([456], dtype=torch.int32))], ) @@ -43,6 +46,12 @@ def _capture_scatter_args(monkeypatch): return scatter_args +def _assert_all_mem_indexes_are_freed(extra_mem_indexes_cpu, expected): + assert len(extra_mem_indexes_cpu) == 1 + assert torch.equal(extra_mem_indexes_cpu[0].mem_indexes_cpu, expected) + assert extra_mem_indexes_cpu[0].free_mask_cpu is None + + def test_backends_initialize_their_own_spec_engine(): args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False) backend = ChunkedPrefillBackend.__new__(ChunkedPrefillBackend) @@ -126,7 +135,7 @@ def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(monkeypatch): assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) assert scatter_args["proposal"].token_ids.shape == (16, 8) assert scatter_args["valid_row_count"] == 8 - assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) + _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([123], dtype=torch.int32)) def test_dp_vanilla_uses_dp_engine_proposer(monkeypatch): @@ -155,7 +164,7 @@ def test_dp_vanilla_uses_dp_engine_proposer(monkeypatch): assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) assert scatter_args["proposal"].token_ids.shape == (16, 8) assert scatter_args["valid_row_count"] == 8 - assert torch.equal(extra_mem, torch.tensor([123], dtype=torch.int32)) + _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([123], dtype=torch.int32)) def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(monkeypatch): @@ -201,7 +210,7 @@ def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(monkeypat assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) assert scatter_args["proposal"].token_ids.shape == (24, 8) - assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) + _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([456], dtype=torch.int32)) def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(monkeypatch): @@ -242,4 +251,4 @@ def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(monkeypatch) assert propose_args["accept_len0"] is None assert propose_args["accept_len1"] is None assert scatter_args["proposal"].token_ids.shape == (24, 8) - assert torch.equal(extra_mem, torch.tensor([456], dtype=torch.int32)) + _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([456], dtype=torch.int32)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index ecba9e705b..fd717ca820 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -109,7 +109,9 @@ def test_overlap_eagle_keeps_fixed_verify_layout(monkeypatch): expected_draft_tokens = torch.tensor([0, 1, 2, 0, 1, 2, 3, 4, 5]) assert torch.equal(proposal.token_ids[:, 1], expected_draft_tokens) assert torch.equal(proposal.token_ids[:, 2], expected_draft_tokens) - assert torch.equal(proposal.extra_mem_indexes_cpu, torch.arange(6, dtype=torch.int32)) + assert len(proposal.extra_mem_indexes_cpu) == 1 + assert torch.equal(proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, torch.arange(6, dtype=torch.int32)) + assert proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None assert torch.equal(model_input0.mem_indexes, torch.tensor([2, 0, 1, 5, 99, 99], dtype=torch.int32)) assert torch.equal(model_input1.mem_indexes, torch.tensor([2, 2, 4, 5, 3, 5], dtype=torch.int32)) @@ -161,7 +163,9 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): assert draft_model.extend_batch_sizes == (6, 6) assert draft_model.decode_batch_sizes == [(2, 2)] assert proposal.token_ids.shape == (9, 3) - assert torch.equal(proposal.extra_mem_indexes_cpu, torch.arange(3, dtype=torch.int32)) + assert len(proposal.extra_mem_indexes_cpu) == 1 + assert torch.equal(proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, torch.arange(3, dtype=torch.int32)) + assert proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None def test_eagle3_maps_draft_token_ids_in_proposer(): diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 7ce4ca3e3b..f619618e75 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -50,7 +50,11 @@ ) from lightllm.server.router.model_infer.mtp_speculative.planner.base import _InferCostMsTable from lightllm.server.router.model_infer.mtp_speculative.proposers import build_spec_proposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + BaseSpecProposer, + MtpMemIndexesToFree, + SpecProposal, +) from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer, DFlashSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer, DSparkSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer @@ -174,7 +178,7 @@ def test_scatter_mtp_next_tokens_consumes_mode_proposal(monkeypatch): ) proposal = DFlashSpecProposal( token_ids=torch.arange(6, dtype=torch.int64).view(3, 2), - extra_mem_indexes_cpu=None, + extra_mem_indexes_cpu=[], schedule_scores=torch.arange(3, dtype=torch.float32).view(3, 1), ) @@ -378,7 +382,13 @@ def test_dynamic_decode_frees_unselected_and_rejected_rows(): backend=backend, proposal=SpecProposal( token_ids=torch.empty((0,), dtype=torch.int64), - extra_mem_indexes_cpu=torch.tensor([20]), + extra_mem_indexes_cpu=[ + MtpMemIndexesToFree( + mem_indexes_cpu=torch.tensor([20, 21]), + free_mask_cpu=torch.tensor([True, False]), + ), + MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([22])), + ], ), model_input=model_input, selected_row_mask_cpu=torch.tensor([1, 0, 1, 0], dtype=torch.bool), @@ -386,7 +396,7 @@ def test_dynamic_decode_frees_unselected_and_rejected_rows(): ) assert len(freed) == 1 - assert freed[0].tolist() == [11, 12, 13, 20] + assert freed[0].tolist() == [11, 12, 13, 20, 22] def test_records_request_mtp_metrics_in_one_pass(): @@ -737,7 +747,7 @@ def test_engine_skips_feedback_for_a_mixed_proposal_batch(): plan=plan, proposal=SpecProposal( token_ids=torch.empty((0,), dtype=torch.int64), - extra_mem_indexes_cpu=None, + extra_mem_indexes_cpu=[], ), req_num=2, accept_lengths_cpu=torch.tensor([1, 2], dtype=torch.int32), @@ -755,7 +765,7 @@ def test_dspark_applies_confidence_capacity_after_two_step_delay(): plan = SpecDecodePlan(origin_batch_size=8, dynamic_batch_size=8, draft_step=3, pre_draft_step=3) proposal = DSparkSpecProposal( token_ids=torch.empty((0,), dtype=torch.int64), - extra_mem_indexes_cpu=None, + extra_mem_indexes_cpu=[], schedule_scores_cpu=torch.from_numpy(confidence_probs), ) engine = SpecEngine.__new__(SpecEngine) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py index 8e738a9b7f..8cb4761989 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -60,6 +60,6 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): [21, 1, 1], [22, 2, 2], ] - assert proposal.extra_mem_indexes_cpu is None + assert proposal.extra_mem_indexes_cpu == [] assert draft_models[0].decode_batch_sizes == [(4, 4)] assert draft_models[1].decode_batch_sizes == [(4, 4)] From 9a00892886acc07da4dbd081b0fe8d141d23cb70 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 07:28:29 +0000 Subject: [PATCH 054/103] refactor: simplify dynamic MTP memory handling --- .../basemodel/triton_kernel/mtp_utils.py | 30 +------ .../mode_backend/chunked_prefill/impl.py | 21 +++-- .../mode_backend/dp_backend/impl.py | 19 ++++- .../model_infer/mtp_speculative/engine.py | 14 ++++ .../mtp_speculative/proposers/base.py | 4 +- .../model_infer/mtp_speculative/utils.py | 33 +------- .../basemodel/triton_kernel/test_mtp_utils.py | 12 +-- .../mtp_speculative/test_planner.py | 82 +++++++++++++++---- 8 files changed, 119 insertions(+), 96 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index 7f05a975c4..8161b22f6c 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -213,8 +213,6 @@ def _fwd_kernel_compact_dynamic_mtp_model_input( out_b_seq_len, b_position_delta, out_b_position_delta, - mem_indexes, - out_mem_indexes, b_shared_seq_len, out_b_shared_seq_len, selected_mask, @@ -222,7 +220,6 @@ def _fwd_kernel_compact_dynamic_mtp_model_input( batch_size, HAS_INPUT_IDS: tl.constexpr, HAS_B_POSITION_DELTA: tl.constexpr, - HAS_MEM_INDEXES: tl.constexpr, HAS_B_SHARED_SEQ_LEN: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): @@ -250,10 +247,6 @@ def _fwd_kernel_compact_dynamic_mtp_model_input( position_delta = tl.load(b_position_delta + offsets, mask=mask, other=0) tl.store(out_b_position_delta + dst_pos, position_delta, mask=write_mask) - if HAS_MEM_INDEXES: - mem_index = tl.load(mem_indexes + offsets, mask=mask, other=0) - tl.store(out_mem_indexes + dst_pos, mem_index, mask=write_mask) - if HAS_B_SHARED_SEQ_LEN: shared_seq_len = tl.load(b_shared_seq_len + offsets, mask=mask, other=0) tl.store(out_b_shared_seq_len + dst_pos, shared_seq_len, mask=write_mask) @@ -428,15 +421,6 @@ def _compact_decode_model_input( device=model_input.b_position_delta.device, ) - out_mem_indexes = None - if model_input.mem_indexes is not None: - assert model_input.mem_indexes.is_cuda - out_mem_indexes = torch.empty( - (dynamic_batch_size,), - dtype=model_input.mem_indexes.dtype, - device=model_input.mem_indexes.device, - ) - dummy_1d = model_input.b_req_idx BLOCK_SIZE = triton.next_power_of_2(old_batch_size) grid = (1,) @@ -451,8 +435,6 @@ def _compact_decode_model_input( out_b_seq_len=out_b_seq_len, b_position_delta=model_input.b_position_delta if model_input.b_position_delta is not None else dummy_1d, out_b_position_delta=out_b_position_delta if out_b_position_delta is not None else dummy_1d, - mem_indexes=model_input.mem_indexes if model_input.mem_indexes is not None else dummy_1d, - out_mem_indexes=out_mem_indexes if out_mem_indexes is not None else dummy_1d, b_shared_seq_len=model_input.b_shared_seq_len if model_input.b_shared_seq_len is not None else dummy_1d, out_b_shared_seq_len=out_b_shared_seq_len if out_b_shared_seq_len is not None else dummy_1d, selected_mask=selected_row_mask, @@ -460,7 +442,6 @@ def _compact_decode_model_input( batch_size=old_batch_size, HAS_INPUT_IDS=model_input.input_ids is not None, HAS_B_POSITION_DELTA=model_input.b_position_delta is not None, - HAS_MEM_INDEXES=model_input.mem_indexes is not None, HAS_B_SHARED_SEQ_LEN=model_input.b_shared_seq_len is not None, BLOCK_SIZE=BLOCK_SIZE, num_warps=8, @@ -472,7 +453,6 @@ def _compact_decode_model_input( model_input.b_mtp_index = out_b_mtp_index model_input.b_seq_len = out_b_seq_len model_input.b_position_delta = out_b_position_delta - model_input.mem_indexes = out_mem_indexes model_input.b_shared_seq_len = out_b_shared_seq_len model_input.b_mark_shared_group = _rebuild_mtp_group_markers( out_b_req_idx, @@ -513,6 +493,7 @@ def prepare_dynamic_mtp_model_input( # All compaction work stays on the current CUDA stream and needs no host sync. model_input.to_cuda() + assert model_input.mem_indexes.shape[0] == dynamic_batch_size selected_row_mask = sample_dynamic_mtp_row_mask( dynamic_batch_size=dynamic_batch_size, @@ -528,12 +509,9 @@ def prepare_dynamic_mtp_model_input( dynamic_batch_size=dynamic_batch_size, max_draft_step=max_draft_step, ) - # Keep CPU mem_indexes unfiltered here. Copying selected_row_mask back to - # CPU in this hot path synchronizes the overlap stream; the router frees - # unselected/rejected CPU mem indexes after its existing async mask copy is - # consumed. Decode and draft-cache commit use the compacted - # b_position_delta, so placeholder multimodal metadata only needs to keep - # ModelInput shapes consistent. + # Decode and draft-cache commit use the compacted b_position_delta, so + # placeholder multimodal metadata only needs to keep ModelInput shapes + # consistent. if model_input.multimodal_params is not None: # Read-only placeholders: avoid rebuilding hundreds of nested Python # objects on every compacted decode iteration. diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 045a5a056c..f80d86a200 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -14,6 +14,7 @@ from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_current_device_id from .control_state import ControlState @@ -267,8 +268,8 @@ def decode_mtp( # 长度和顺序与压缩后的 model_output.logits 保持一一对应,供后续采样使用。 if async_selected_row_mask_cpu is not None: async_selected_row_mask_cpu.wait() - selected_row_mask_cpu = async_selected_row_mask_cpu.tensor.tolist() - run_reqs = [req for req, selected in zip(run_reqs, selected_row_mask_cpu) if selected] + selected_rows = async_selected_row_mask_cpu.tensor.tolist() + run_reqs = [req for req, selected in zip(run_reqs, selected_rows) if selected] next_token_ids, next_token_logprobs = sample( model_output.logits, run_reqs, @@ -374,14 +375,16 @@ def decode_mtp( extra_post_req_handle_func=self.extra_post_req_handle_func, ) - mtp_utils.free_unused_mtp_decode_mem( - backend=self, - proposal=proposal, - model_input=model_input, - selected_row_mask_cpu=( - async_selected_row_mask_cpu.tensor if async_selected_row_mask_cpu is not None else None + proposal.extra_mem_indexes_cpu.insert( + 0, + MtpMemIndexesToFree( + mem_indexes_cpu=model_input.mem_indexes_cpu, + free_mask_cpu=accepted_index_cpu == 0, ), - accepted_index_cpu=accepted_index_cpu, + ) + mtp_utils.free_mem_indexes( + backend=self, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, ) # 第四阶段 diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 8baf032a7a..7de0541cd4 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -18,6 +18,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from .control_state import DPControlState @@ -600,7 +601,13 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): # 第三阶段 event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() - need_free_mem_indexes = model_input.mem_indexes_cpu[0:req_num][accepted_index_cpu == 0] + extra_mem_indexes_cpu.insert( + 0, + MtpMemIndexesToFree( + mem_indexes_cpu=model_input.mem_indexes_cpu[0:req_num], + free_mask_cpu=accepted_index_cpu == 0, + ), + ) select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( @@ -613,7 +620,6 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): ) mtp_utils.free_mem_indexes( backend=self, - mem_indexes_cpu=need_free_mem_indexes, extra_mem_indexes_cpu=extra_mem_indexes_cpu, ) @@ -924,7 +930,13 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf mem_indexes_cpu = torch.cat( (model_input0.mem_indexes_cpu[0:req_num0], model_input1.mem_indexes_cpu[0:req_num1]), dim=0 ) - need_free_mem_indexes = mem_indexes_cpu[accepted_index_cpu == 0] + extra_mem_indexes_cpu.insert( + 0, + MtpMemIndexesToFree( + mem_indexes_cpu=mem_indexes_cpu, + free_mask_cpu=accepted_index_cpu == 0, + ), + ) select_mask = accepted_index_cpu.to(dtype=torch.bool) self._post_handle( @@ -937,7 +949,6 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf ) mtp_utils.free_mem_indexes( backend=self, - mem_indexes_cpu=need_free_mem_indexes, extra_mem_indexes_cpu=extra_mem_indexes_cpu, ) event_pack.notify_pre_post_handle() diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 107fed3e1e..474fc9c9d8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -77,6 +77,20 @@ def prepare_decode_model_input( return model_input, None from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_mtp_model_input + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + # mem_indexes 是本轮 decode 新申请、尚未绑定请求和 token 位置的 KV slot。 + # 动态 verify 只需要保留 dynamic_batch_size 个任意 slot,因此 CPU 和已存在的 + # GPU 索引都可以直接截取前缀,无需等待 selected_row_mask_cpu。多申请的 CPU + # 尾部索引在这里立即归还;后续 forward 会根据压缩后的 b_req_idx/b_seq_len + # 建立保留 slot 与实际请求位置之间的映射。该操作需要放在下方动态输入构建 + # 之前,避免其内部 to_cuda 将原始完整 batch 的 mem indexes 全量复制到 GPU。 + unused_mem_indexes_cpu = model_input.mem_indexes_cpu[plan.dynamic_batch_size :] + model_input.mem_indexes_cpu = model_input.mem_indexes_cpu[: plan.dynamic_batch_size] + if model_input.mem_indexes is not None: + model_input.mem_indexes = model_input.mem_indexes[: plan.dynamic_batch_size] + if unused_mem_indexes_cpu.numel() > 0: + g_infer_context.req_manager.mem_manager.free(unused_mem_indexes_cpu) model_input, selected_row_mask = prepare_dynamic_mtp_model_input( model_input=model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index b8e16fc030..065149bfb5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -36,8 +36,8 @@ class SpecProposal: `token_ids` has shape `[verify_batch, draft_step + 1]`; column 0 contains target-model tokens and the remaining columns contain draft candidates. - `extra_mem_indexes_cpu` tracks temporary KV slots owned by the proposal; - each item describes one independently managed group of temporary indexes. + `extra_mem_indexes_cpu` uniformly tracks every KV slot considered for + release, including rejected target rows and proposal-owned temporary rows. Mode-specific scheduling metadata belongs to the corresponding subclass. """ diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index 01789171d0..c6cfa48858 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -5,7 +5,6 @@ import torch -from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.basemodel.triton_kernel.mtp_utils import ( linear_att_mtp_state_index_update, mtp_scatter_next_token_ids, @@ -124,15 +123,11 @@ def record_request_mtp_metrics( def free_mem_indexes( backend: ModeBackend, - mem_indexes_cpu: torch.Tensor, extra_mem_indexes_cpu: List[MtpMemIndexesToFree], ) -> None: - """Free regular KV indexes plus proposal-owned indexes selected by masks.""" + """Free all KV indexes described by the unified MTP memory list.""" mem_indexes_to_free = [] - if mem_indexes_cpu.numel() > 0: - mem_indexes_to_free.append(mem_indexes_cpu) - for extra_mem_to_free in extra_mem_indexes_cpu: extra_indexes_cpu = extra_mem_to_free.mem_indexes_cpu if extra_mem_to_free.free_mask_cpu is not None: @@ -144,35 +139,9 @@ def free_mem_indexes( backend.model.req_manager.mem_manager.free(torch.cat(mem_indexes_to_free, dim=0)) -def free_unused_mtp_decode_mem( - backend: ModeBackend, - proposal: SpecProposal, - model_input: ModelInput, - selected_row_mask_cpu: Optional[torch.Tensor], - accepted_index_cpu: torch.Tensor, -) -> None: - """Free rejected target KV slots and draft-only temporary slots.""" - - mem_indexes_cpu = model_input.mem_indexes_cpu - if selected_row_mask_cpu is None: - free_mask = accepted_index_cpu == 0 - else: - selected_mask = selected_row_mask_cpu.to(dtype=torch.bool) - free_mask = selected_mask.logical_not() - free_mask[selected_mask] = accepted_index_cpu == 0 - need_free_mem_indexes = mem_indexes_cpu[free_mask] - - free_mem_indexes( - backend=backend, - mem_indexes_cpu=need_free_mem_indexes, - extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, - ) - - __all__ = [ "alloc_mem_indexes", "free_mem_indexes", - "free_unused_mtp_decode_mem", "record_request_mtp_metrics", "scatter_mtp_next_tokens", "verify_mtp_tokens", diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py index faefb2863c..2e1669570a 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -48,8 +48,8 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): b_position_delta=torch.arange(12, dtype=torch.int32, device="cuda") + 200, b_shared_seq_len=torch.tensor([0, 0, 0, 0, 7, 7, 7, 7, 9, 9, 9, 9], dtype=torch.int32, device="cuda"), b_mark_shared_group=torch.tensor([0, 0, 0, 4, 0, 0, 0, 4, 0, 0, 0, 4], dtype=torch.int32, device="cuda"), - mem_indexes=torch.arange(12, dtype=torch.int32, device="cuda") + 100, - mem_indexes_cpu=torch.arange(12, dtype=torch.int32, device="cpu") + 100, + mem_indexes=torch.arange(8, dtype=torch.int32, device="cuda") + 100, + mem_indexes_cpu=torch.arange(8, dtype=torch.int32, device="cpu") + 100, is_prefill=False, multimodal_params=[{"row": i, "images": [], "audios": []} for i in range(12)], mtp_draft_input_hiddens=(torch.arange(12 * 5, dtype=torch.float32, device="cuda").reshape(12, 5) + 0.5), @@ -97,11 +97,11 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 1, 0, 0, 0, 4], dtype=torch.int32) ) assert torch.equal( - compacted_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 104, 108, 109, 110, 111], dtype=torch.int32) + compacted_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 103, 104, 105, 106, 107], dtype=torch.int32) ) - # The hot path intentionally keeps the CPU copy unfiltered to avoid a GPU-to-CPU - # synchronization. The router frees rejected indexes after its async mask copy. - assert torch.equal(compacted_input.mem_indexes_cpu, torch.arange(12, dtype=torch.int32) + 100) + # CPU/GPU mem indexes are trimmed by SpecEngine.prepare_decode_model_input + # before entering this lower-level row compaction helper. + assert torch.equal(compacted_input.mem_indexes_cpu, torch.arange(8, dtype=torch.int32) + 100) expected_hiddens = (torch.arange(12 * 5, dtype=torch.float32).reshape(12, 5) + 0.5)[expected_selected_rows] assert torch.equal(compacted_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index f619618e75..7b7e51f6f6 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -369,30 +369,23 @@ def test_eagle_proposer_skips_draft_forward_for_zero_steps(): assert proposal.schedule_scores.shape == (2, 0) -def test_dynamic_decode_frees_unselected_and_rejected_rows(): +def test_free_mem_indexes_applies_extra_masks(): freed = [] backend = SimpleNamespace( model=SimpleNamespace( req_manager=SimpleNamespace(mem_manager=SimpleNamespace(free=lambda indexes: freed.append(indexes.clone()))) ) ) - model_input = SimpleNamespace(mem_indexes_cpu=torch.tensor([10, 11, 12, 13])) - - mtp_utils.free_unused_mtp_decode_mem( + mtp_utils.free_mem_indexes( backend=backend, - proposal=SpecProposal( - token_ids=torch.empty((0,), dtype=torch.int64), - extra_mem_indexes_cpu=[ - MtpMemIndexesToFree( - mem_indexes_cpu=torch.tensor([20, 21]), - free_mask_cpu=torch.tensor([True, False]), - ), - MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([22])), - ], - ), - model_input=model_input, - selected_row_mask_cpu=torch.tensor([1, 0, 1, 0], dtype=torch.bool), - accepted_index_cpu=torch.tensor([1, 0], dtype=torch.int32), + extra_mem_indexes_cpu=[ + MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([11, 12, 13])), + MtpMemIndexesToFree( + mem_indexes_cpu=torch.tensor([20, 21]), + free_mask_cpu=torch.tensor([True, False]), + ), + MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([22])), + ], ) assert len(freed) == 1 @@ -543,6 +536,61 @@ def test_engine_lets_planner_count_requests_with_a_previous_proposal(): assert ready_plan.all_reqs_have_proposals +def test_dynamic_prepare_keeps_prefix_mem_indexes_and_frees_unused_tail(monkeypatch): + from lightllm.common.basemodel.triton_kernel import mtp_utils as common_mtp_utils + from lightllm.server.router.model_infer import infer_batch as infer_batch_module + from lightllm.server.router.model_infer.mtp_speculative import engine as engine_module + + freed = [] + monkeypatch.setattr( + infer_batch_module, + "g_infer_context", + SimpleNamespace( + req_manager=SimpleNamespace(mem_manager=SimpleNamespace(free=lambda indexes: freed.append(indexes.clone()))) + ), + ) + selected_row_mask = torch.tensor([1, 0, 1, 0], dtype=torch.int32) + monkeypatch.setattr( + common_mtp_utils, + "prepare_dynamic_mtp_model_input", + lambda model_input, **kwargs: (model_input, selected_row_mask), + ) + async_mask = object() + monkeypatch.setattr( + engine_module.g_pin_mem_manager, + "async_copy_from_gpu_tensor_with_event", + lambda **kwargs: async_mask, + ) + + engine = SpecEngine.__new__(SpecEngine) + engine.backend = SimpleNamespace( + model=SimpleNamespace( + req_manager=SimpleNamespace( + req_sampling_params_manager=SimpleNamespace(req_to_next_token_scores=torch.empty(0)) + ) + ) + ) + model_input = SimpleNamespace( + batch_size=4, + mem_indexes=torch.tensor([20, 21, 22, 23], dtype=torch.int32), + mem_indexes_cpu=torch.tensor([10, 11, 12, 13], dtype=torch.int32), + ) + plan = SpecDecodePlan(origin_batch_size=4, dynamic_batch_size=2, draft_step=1, pre_draft_step=1) + + compacted_input, selected_mask_cpu = engine.prepare_decode_model_input( + model_input=model_input, + req_num=1, + plan=plan, + ) + + assert compacted_input is model_input + assert selected_mask_cpu is async_mask + assert model_input.mem_indexes.tolist() == [20, 21] + assert model_input.mem_indexes_cpu.tolist() == [10, 11] + assert len(freed) == 1 + assert freed[0].tolist() == [12, 13] + + def test_lightspec_eagle_draft_always_keeps_the_extend_candidate(): planner = build_lightspec_planner(spec_mode="eagle_with_att") for batch_size in (2, 4, 8): From b135abefdbe32ba5eb684c9a796bcde7f8f8c7d5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 07:49:29 +0000 Subject: [PATCH 055/103] refactor: separate dynamic MTP kernels --- .../triton_kernel/dynamic_mtp_utils.py | 331 ++++++++++++++++++ .../basemodel/triton_kernel/mtp_utils.py | 329 +---------------- .../model_infer/mtp_speculative/engine.py | 2 +- .../triton_kernel/test_dynamic_mtp_utils.py | 137 ++++++++ .../basemodel/triton_kernel/test_mtp_utils.py | 135 ------- .../mtp_speculative/test_planner.py | 4 +- 6 files changed, 474 insertions(+), 464 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py index b935e7c8a7..70a5eaf700 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -1,8 +1,17 @@ +"""仅动态 MTP verify 使用的选行与 ModelInput 压缩算子。""" + +from typing import Optional + import triton import triton.language as tl import torch +from lightllm.common.basemodel.batch_objs import ModelInput +from lightllm.common.basemodel.mtp_manager import MtpManager +from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size + +# 动态 verify 行选择。 @triton.jit def _fwd_kernel_cumprod_scores( req_to_next_token_scores, @@ -73,3 +82,325 @@ def sample_dynamic_mtp_row_mask( selected_row_mask = torch.zeros((len(b_req_idx),), dtype=torch.int32, device=b_req_idx.device) selected_row_mask.scatter_(0, selected_ids, 1) return selected_row_mask + + +# 动态 ModelInput 行压缩。 +@triton.jit +def _fwd_kernel_compact_dynamic_mtp_model_input( + input_ids, + out_input_ids, + b_req_idx, + out_b_req_idx, + b_mtp_index, + out_b_mtp_index, + b_seq_len, + out_b_seq_len, + b_position_delta, + out_b_position_delta, + b_shared_seq_len, + out_b_shared_seq_len, + selected_mask, + selected_dst_pos, + batch_size, + HAS_INPUT_IDS: tl.constexpr, + HAS_B_POSITION_DELTA: tl.constexpr, + HAS_B_SHARED_SEQ_LEN: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size + selected_i32 = tl.load(selected_mask + offsets, mask=mask, other=0) + selected = selected_i32 != 0 + dst_pos = tl.cumsum(selected_i32, axis=0) - 1 + write_mask = mask & selected + + cur_b_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=0) + cur_b_mtp_index = tl.load(b_mtp_index + offsets, mask=mask, other=0) + cur_b_seq_len = tl.load(b_seq_len + offsets, mask=mask, other=0) + + tl.store(selected_dst_pos + offsets, dst_pos, mask=mask) + tl.store(out_b_req_idx + dst_pos, cur_b_req_idx, mask=write_mask) + tl.store(out_b_mtp_index + dst_pos, cur_b_mtp_index, mask=write_mask) + tl.store(out_b_seq_len + dst_pos, cur_b_seq_len, mask=write_mask) + + if HAS_INPUT_IDS: + input_id = tl.load(input_ids + offsets, mask=mask, other=0) + tl.store(out_input_ids + dst_pos, input_id, mask=write_mask) + + if HAS_B_POSITION_DELTA: + position_delta = tl.load(b_position_delta + offsets, mask=mask, other=0) + tl.store(out_b_position_delta + dst_pos, position_delta, mask=write_mask) + + if HAS_B_SHARED_SEQ_LEN: + shared_seq_len = tl.load(b_shared_seq_len + offsets, mask=mask, other=0) + tl.store(out_b_shared_seq_len + dst_pos, shared_seq_len, mask=write_mask) + + return + + +@triton.jit +def _fwd_kernel_rebuild_trimmed_mtp_b_mark_shared_group( + b_req_idx, + out_b_mark_shared_group, + batch_size, + max_batch_shared_group_size: tl.constexpr, + MAX_RUN_SCAN: tl.constexpr, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.arange(0, BLOCK_SIZE) + mask = offsets < batch_size + cur_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=-1) + + prev_same_count = tl.full((BLOCK_SIZE,), 0, tl.int32) + for scan_offset in tl.static_range(1, MAX_RUN_SCAN + 1): + prev_offsets = offsets - scan_offset + prev_mask = mask & (prev_offsets >= 0) + prev_req_idx = tl.load(b_req_idx + prev_offsets, mask=prev_mask, other=-2) + prev_same_count += tl.where(prev_mask & (prev_req_idx == cur_req_idx), 1, 0) + + next_offsets = offsets + 1 + next_req_idx = tl.load(b_req_idx + next_offsets, mask=next_offsets < batch_size, other=-2) + group_pos = prev_same_count % max_batch_shared_group_size + is_group_end = mask & ( + (next_offsets == batch_size) | (next_req_idx != cur_req_idx) | (group_pos == max_batch_shared_group_size - 1) + ) + mark_value = tl.where(is_group_end, group_pos + 1, 0) + tl.store(out_b_mark_shared_group + offsets, mark_value, mask=mask) + return + + +@triton.jit +def _fwd_kernel_pack_selected_rows_2d( + src, + src_stride_0, + src_stride_1, + dst, + dst_stride_0, + dst_stride_1, + selected_mask, + selected_dst_pos, + batch_size, + hidden_size, + BLOCK_N: tl.constexpr, +): + row_id = tl.program_id(0) + col_block_id = tl.program_id(1) + col_offsets = col_block_id * BLOCK_N + tl.arange(0, BLOCK_N) + row_mask = row_id < batch_size + + selected_val = tl.load(selected_mask + row_id, mask=row_mask, other=0) + dst_row = tl.load(selected_dst_pos + row_id, mask=row_mask, other=0) + write_row = row_mask & (selected_val != 0) + col_mask = col_offsets < hidden_size + src_ptrs = src + row_id * src_stride_0 + col_offsets * src_stride_1 + dst_ptrs = dst + dst_row * dst_stride_0 + col_offsets * dst_stride_1 + vals = tl.load(src_ptrs, mask=write_row & col_mask, other=0) + tl.store(dst_ptrs, vals, mask=write_row & col_mask) + + +def _pack_selected_hidden( + hidden: torch.Tensor, + selected_row_mask: torch.Tensor, + selected_dst_pos: torch.Tensor, + dynamic_batch_size: int, +): + assert hidden.is_cuda + assert hidden.ndim == 2 + assert selected_row_mask.is_cuda + assert selected_dst_pos.is_cuda + assert hidden.shape[0] == selected_row_mask.shape[0] + + selected_row_mask = selected_row_mask.to(torch.int32) + hidden_size = hidden.shape[1] + dst = torch.empty((dynamic_batch_size, hidden_size), dtype=hidden.dtype, device=hidden.device) + grid = (hidden.shape[0], triton.cdiv(hidden_size, 128)) + _fwd_kernel_pack_selected_rows_2d[grid]( + src=hidden, + src_stride_0=hidden.stride(0), + src_stride_1=hidden.stride(1), + dst=dst, + dst_stride_0=dst.stride(0), + dst_stride_1=dst.stride(1), + selected_mask=selected_row_mask, + selected_dst_pos=selected_dst_pos, + batch_size=hidden.shape[0], + hidden_size=hidden_size, + BLOCK_N=128, + num_warps=4, + num_stages=1, + ) + return dst + + +def _rebuild_mtp_group_markers(b_req_idx: torch.Tensor, max_request_rows: int) -> torch.Tensor: + assert b_req_idx.is_cuda + batch_size = b_req_idx.shape[0] + max_batch_shared_group_size = int(get_diverse_max_batch_shared_group_size()) + assert max_batch_shared_group_size > 0 + assert max_request_rows > 0 + if batch_size == 0: + return torch.empty((0,), dtype=torch.int32, device=b_req_idx.device) + + b_mark_shared_group = torch.empty((batch_size,), dtype=torch.int32, device=b_req_idx.device) + BLOCK_SIZE = triton.next_power_of_2(batch_size) + _fwd_kernel_rebuild_trimmed_mtp_b_mark_shared_group[(1,)]( + b_req_idx=b_req_idx, + out_b_mark_shared_group=b_mark_shared_group, + batch_size=batch_size, + max_batch_shared_group_size=max_batch_shared_group_size, + MAX_RUN_SCAN=max_request_rows - 1, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=8, + num_stages=1, + ) + return b_mark_shared_group + + +def _compact_decode_model_input( + model_input: ModelInput, + selected_row_mask: torch.Tensor, + dynamic_batch_size: int, + max_draft_step: int, +) -> ModelInput: + assert not model_input.is_prefill + assert selected_row_mask.is_cuda + assert model_input.b_req_idx.is_cuda + assert model_input.b_mtp_index.is_cuda + assert model_input.b_seq_len.is_cuda + + # Dynamic scheduling guarantees exactly dynamic_batch_size selected rows. + selected_row_mask = selected_row_mask.to(torch.int32) + old_batch_size = model_input.b_req_idx.shape[0] + selected_dst_pos = torch.empty((old_batch_size,), dtype=torch.int32, device=model_input.b_req_idx.device) + + out_input_ids = None + if model_input.input_ids is not None: + assert model_input.input_ids.is_cuda + out_input_ids = torch.empty( + (dynamic_batch_size,), dtype=model_input.input_ids.dtype, device=model_input.input_ids.device + ) + + out_b_shared_seq_len = None + if model_input.b_shared_seq_len is not None: + assert model_input.b_shared_seq_len.is_cuda + out_b_shared_seq_len = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_shared_seq_len.dtype, device=model_input.b_shared_seq_len.device + ) + + out_b_req_idx = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_req_idx.dtype, device=model_input.b_req_idx.device + ) + out_b_mtp_index = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_mtp_index.dtype, device=model_input.b_mtp_index.device + ) + out_b_seq_len = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_seq_len.dtype, device=model_input.b_seq_len.device + ) + out_b_position_delta = None + if model_input.b_position_delta is not None: + assert model_input.b_position_delta.is_cuda + out_b_position_delta = torch.empty( + (dynamic_batch_size,), + dtype=model_input.b_position_delta.dtype, + device=model_input.b_position_delta.device, + ) + + dummy_1d = model_input.b_req_idx + BLOCK_SIZE = triton.next_power_of_2(old_batch_size) + grid = (1,) + _fwd_kernel_compact_dynamic_mtp_model_input[grid]( + input_ids=model_input.input_ids if model_input.input_ids is not None else dummy_1d, + out_input_ids=out_input_ids if out_input_ids is not None else dummy_1d, + b_req_idx=model_input.b_req_idx, + out_b_req_idx=out_b_req_idx, + b_mtp_index=model_input.b_mtp_index, + out_b_mtp_index=out_b_mtp_index, + b_seq_len=model_input.b_seq_len, + out_b_seq_len=out_b_seq_len, + b_position_delta=model_input.b_position_delta if model_input.b_position_delta is not None else dummy_1d, + out_b_position_delta=out_b_position_delta if out_b_position_delta is not None else dummy_1d, + b_shared_seq_len=model_input.b_shared_seq_len if model_input.b_shared_seq_len is not None else dummy_1d, + out_b_shared_seq_len=out_b_shared_seq_len if out_b_shared_seq_len is not None else dummy_1d, + selected_mask=selected_row_mask, + selected_dst_pos=selected_dst_pos, + batch_size=old_batch_size, + HAS_INPUT_IDS=model_input.input_ids is not None, + HAS_B_POSITION_DELTA=model_input.b_position_delta is not None, + HAS_B_SHARED_SEQ_LEN=model_input.b_shared_seq_len is not None, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=8, + num_stages=1, + ) + + model_input.input_ids = out_input_ids + model_input.b_req_idx = out_b_req_idx + model_input.b_mtp_index = out_b_mtp_index + model_input.b_seq_len = out_b_seq_len + model_input.b_position_delta = out_b_position_delta + model_input.b_shared_seq_len = out_b_shared_seq_len + model_input.b_mark_shared_group = _rebuild_mtp_group_markers( + out_b_req_idx, + max_request_rows=max_draft_step + 1, + ) + + if model_input.mtp_draft_input_hiddens is not None: + assert model_input.mtp_draft_input_hiddens.is_cuda + model_input.mtp_draft_input_hiddens = _pack_selected_hidden( + model_input.mtp_draft_input_hiddens, + selected_row_mask, + selected_dst_pos, + dynamic_batch_size, + ) + model_input.batch_size = dynamic_batch_size + + return model_input + + +def prepare_dynamic_mtp_model_input( + model_input: ModelInput, + req_num: int, + dynamic_batch_size: int, + req_to_next_token_scores: torch.Tensor, + pre_draft_step: Optional[int] = None, +): + req_num = int(req_num) + dynamic_batch_size = int(dynamic_batch_size) + assert not model_input.is_prefill, "prepare_dynamic_mtp_model_input only supports decode inputs" + assert req_to_next_token_scores is not None + assert dynamic_batch_size >= req_num + assert dynamic_batch_size <= model_input.batch_size + max_draft_step = MtpManager.get_instance().get_decode_draft_step(is_draft_model=False) + pre_draft_step = max_draft_step if pre_draft_step is None else int(pre_draft_step) + assert 0 <= pre_draft_step <= max_draft_step + assert model_input.batch_size == req_num * (max_draft_step + 1) + assert dynamic_batch_size <= req_num * (pre_draft_step + 1) + + # All compaction work stays on the current CUDA stream and needs no host sync. + model_input.to_cuda() + assert model_input.mem_indexes.shape[0] == dynamic_batch_size + + selected_row_mask = sample_dynamic_mtp_row_mask( + dynamic_batch_size=dynamic_batch_size, + b_req_idx=model_input.b_req_idx, + req_to_next_token_scores=req_to_next_token_scores, + max_draft_step=max_draft_step, + pre_draft_step=pre_draft_step, + ) + + model_input = _compact_decode_model_input( + model_input=model_input, + selected_row_mask=selected_row_mask, + dynamic_batch_size=dynamic_batch_size, + max_draft_step=max_draft_step, + ) + # Decode and draft-cache commit use the compacted b_position_delta, so + # placeholder multimodal metadata only needs to keep ModelInput shapes + # consistent. + if model_input.multimodal_params is not None: + # Read-only placeholders: avoid rebuilding hundreds of nested Python + # objects on every compacted decode iteration. + empty_multimodal_params = {"images": [], "audios": []} + model_input.multimodal_params = [empty_multimodal_params] * dynamic_batch_size + + model_input.max_q_seq_len = 1 + return model_input, selected_row_mask diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index 8161b22f6c..cb13fad7f4 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -1,13 +1,11 @@ +"""固定布局与动态布局共同使用的 MTP Triton 算子。""" + from typing import Optional + import triton import triton.language as tl import torch -from lightllm.common.basemodel.batch_objs import ModelInput -from lightllm.common.basemodel.mtp_manager import MtpManager -from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import sample_dynamic_mtp_row_mask -from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size - @triton.jit def _fwd_kernel_mtp_verify( @@ -201,327 +199,6 @@ def mtp_scatter_next_token_ids( ) -@triton.jit -def _fwd_kernel_compact_dynamic_mtp_model_input( - input_ids, - out_input_ids, - b_req_idx, - out_b_req_idx, - b_mtp_index, - out_b_mtp_index, - b_seq_len, - out_b_seq_len, - b_position_delta, - out_b_position_delta, - b_shared_seq_len, - out_b_shared_seq_len, - selected_mask, - selected_dst_pos, - batch_size, - HAS_INPUT_IDS: tl.constexpr, - HAS_B_POSITION_DELTA: tl.constexpr, - HAS_B_SHARED_SEQ_LEN: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - offsets = tl.arange(0, BLOCK_SIZE) - mask = offsets < batch_size - selected_i32 = tl.load(selected_mask + offsets, mask=mask, other=0) - selected = selected_i32 != 0 - dst_pos = tl.cumsum(selected_i32, axis=0) - 1 - write_mask = mask & selected - - cur_b_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=0) - cur_b_mtp_index = tl.load(b_mtp_index + offsets, mask=mask, other=0) - cur_b_seq_len = tl.load(b_seq_len + offsets, mask=mask, other=0) - - tl.store(selected_dst_pos + offsets, dst_pos, mask=mask) - tl.store(out_b_req_idx + dst_pos, cur_b_req_idx, mask=write_mask) - tl.store(out_b_mtp_index + dst_pos, cur_b_mtp_index, mask=write_mask) - tl.store(out_b_seq_len + dst_pos, cur_b_seq_len, mask=write_mask) - - if HAS_INPUT_IDS: - input_id = tl.load(input_ids + offsets, mask=mask, other=0) - tl.store(out_input_ids + dst_pos, input_id, mask=write_mask) - - if HAS_B_POSITION_DELTA: - position_delta = tl.load(b_position_delta + offsets, mask=mask, other=0) - tl.store(out_b_position_delta + dst_pos, position_delta, mask=write_mask) - - if HAS_B_SHARED_SEQ_LEN: - shared_seq_len = tl.load(b_shared_seq_len + offsets, mask=mask, other=0) - tl.store(out_b_shared_seq_len + dst_pos, shared_seq_len, mask=write_mask) - - return - - -@triton.jit -def _fwd_kernel_rebuild_trimmed_mtp_b_mark_shared_group( - b_req_idx, - out_b_mark_shared_group, - batch_size, - max_batch_shared_group_size: tl.constexpr, - MAX_RUN_SCAN: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - offsets = tl.arange(0, BLOCK_SIZE) - mask = offsets < batch_size - cur_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=-1) - - prev_same_count = tl.full((BLOCK_SIZE,), 0, tl.int32) - for scan_offset in tl.static_range(1, MAX_RUN_SCAN + 1): - prev_offsets = offsets - scan_offset - prev_mask = mask & (prev_offsets >= 0) - prev_req_idx = tl.load(b_req_idx + prev_offsets, mask=prev_mask, other=-2) - prev_same_count += tl.where(prev_mask & (prev_req_idx == cur_req_idx), 1, 0) - - next_offsets = offsets + 1 - next_req_idx = tl.load(b_req_idx + next_offsets, mask=next_offsets < batch_size, other=-2) - group_pos = prev_same_count % max_batch_shared_group_size - is_group_end = mask & ( - (next_offsets == batch_size) | (next_req_idx != cur_req_idx) | (group_pos == max_batch_shared_group_size - 1) - ) - mark_value = tl.where(is_group_end, group_pos + 1, 0) - tl.store(out_b_mark_shared_group + offsets, mark_value, mask=mask) - return - - -@triton.jit -def _fwd_kernel_pack_selected_rows_2d( - src, - src_stride_0, - src_stride_1, - dst, - dst_stride_0, - dst_stride_1, - selected_mask, - selected_dst_pos, - batch_size, - hidden_size, - BLOCK_N: tl.constexpr, -): - row_id = tl.program_id(0) - col_block_id = tl.program_id(1) - col_offsets = col_block_id * BLOCK_N + tl.arange(0, BLOCK_N) - row_mask = row_id < batch_size - - selected_val = tl.load(selected_mask + row_id, mask=row_mask, other=0) - dst_row = tl.load(selected_dst_pos + row_id, mask=row_mask, other=0) - write_row = row_mask & (selected_val != 0) - col_mask = col_offsets < hidden_size - src_ptrs = src + row_id * src_stride_0 + col_offsets * src_stride_1 - dst_ptrs = dst + dst_row * dst_stride_0 + col_offsets * dst_stride_1 - vals = tl.load(src_ptrs, mask=write_row & col_mask, other=0) - tl.store(dst_ptrs, vals, mask=write_row & col_mask) - - -def _pack_selected_hidden( - hidden: torch.Tensor, - selected_row_mask: torch.Tensor, - selected_dst_pos: torch.Tensor, - dynamic_batch_size: int, -): - assert hidden.is_cuda - assert hidden.ndim == 2 - assert selected_row_mask.is_cuda - assert selected_dst_pos.is_cuda - assert hidden.shape[0] == selected_row_mask.shape[0] - - selected_row_mask = selected_row_mask.to(torch.int32) - hidden_size = hidden.shape[1] - dst = torch.empty((dynamic_batch_size, hidden_size), dtype=hidden.dtype, device=hidden.device) - grid = (hidden.shape[0], triton.cdiv(hidden_size, 128)) - _fwd_kernel_pack_selected_rows_2d[grid]( - src=hidden, - src_stride_0=hidden.stride(0), - src_stride_1=hidden.stride(1), - dst=dst, - dst_stride_0=dst.stride(0), - dst_stride_1=dst.stride(1), - selected_mask=selected_row_mask, - selected_dst_pos=selected_dst_pos, - batch_size=hidden.shape[0], - hidden_size=hidden_size, - BLOCK_N=128, - num_warps=4, - num_stages=1, - ) - return dst - - -def _rebuild_mtp_group_markers(b_req_idx: torch.Tensor, max_request_rows: int) -> torch.Tensor: - assert b_req_idx.is_cuda - batch_size = b_req_idx.shape[0] - max_batch_shared_group_size = int(get_diverse_max_batch_shared_group_size()) - assert max_batch_shared_group_size > 0 - assert max_request_rows > 0 - if batch_size == 0: - return torch.empty((0,), dtype=torch.int32, device=b_req_idx.device) - - b_mark_shared_group = torch.empty((batch_size,), dtype=torch.int32, device=b_req_idx.device) - BLOCK_SIZE = triton.next_power_of_2(batch_size) - _fwd_kernel_rebuild_trimmed_mtp_b_mark_shared_group[(1,)]( - b_req_idx=b_req_idx, - out_b_mark_shared_group=b_mark_shared_group, - batch_size=batch_size, - max_batch_shared_group_size=max_batch_shared_group_size, - MAX_RUN_SCAN=max_request_rows - 1, - BLOCK_SIZE=BLOCK_SIZE, - num_warps=8, - num_stages=1, - ) - return b_mark_shared_group - - -def _compact_decode_model_input( - model_input: ModelInput, - selected_row_mask: torch.Tensor, - dynamic_batch_size: int, - max_draft_step: int, -) -> ModelInput: - assert not model_input.is_prefill - assert selected_row_mask.is_cuda - assert model_input.b_req_idx.is_cuda - assert model_input.b_mtp_index.is_cuda - assert model_input.b_seq_len.is_cuda - - # Dynamic scheduling guarantees exactly dynamic_batch_size selected rows. - selected_row_mask = selected_row_mask.to(torch.int32) - old_batch_size = model_input.b_req_idx.shape[0] - selected_dst_pos = torch.empty((old_batch_size,), dtype=torch.int32, device=model_input.b_req_idx.device) - - out_input_ids = None - if model_input.input_ids is not None: - assert model_input.input_ids.is_cuda - out_input_ids = torch.empty( - (dynamic_batch_size,), dtype=model_input.input_ids.dtype, device=model_input.input_ids.device - ) - - out_b_shared_seq_len = None - if model_input.b_shared_seq_len is not None: - assert model_input.b_shared_seq_len.is_cuda - out_b_shared_seq_len = torch.empty( - (dynamic_batch_size,), dtype=model_input.b_shared_seq_len.dtype, device=model_input.b_shared_seq_len.device - ) - - out_b_req_idx = torch.empty( - (dynamic_batch_size,), dtype=model_input.b_req_idx.dtype, device=model_input.b_req_idx.device - ) - out_b_mtp_index = torch.empty( - (dynamic_batch_size,), dtype=model_input.b_mtp_index.dtype, device=model_input.b_mtp_index.device - ) - out_b_seq_len = torch.empty( - (dynamic_batch_size,), dtype=model_input.b_seq_len.dtype, device=model_input.b_seq_len.device - ) - out_b_position_delta = None - if model_input.b_position_delta is not None: - assert model_input.b_position_delta.is_cuda - out_b_position_delta = torch.empty( - (dynamic_batch_size,), - dtype=model_input.b_position_delta.dtype, - device=model_input.b_position_delta.device, - ) - - dummy_1d = model_input.b_req_idx - BLOCK_SIZE = triton.next_power_of_2(old_batch_size) - grid = (1,) - _fwd_kernel_compact_dynamic_mtp_model_input[grid]( - input_ids=model_input.input_ids if model_input.input_ids is not None else dummy_1d, - out_input_ids=out_input_ids if out_input_ids is not None else dummy_1d, - b_req_idx=model_input.b_req_idx, - out_b_req_idx=out_b_req_idx, - b_mtp_index=model_input.b_mtp_index, - out_b_mtp_index=out_b_mtp_index, - b_seq_len=model_input.b_seq_len, - out_b_seq_len=out_b_seq_len, - b_position_delta=model_input.b_position_delta if model_input.b_position_delta is not None else dummy_1d, - out_b_position_delta=out_b_position_delta if out_b_position_delta is not None else dummy_1d, - b_shared_seq_len=model_input.b_shared_seq_len if model_input.b_shared_seq_len is not None else dummy_1d, - out_b_shared_seq_len=out_b_shared_seq_len if out_b_shared_seq_len is not None else dummy_1d, - selected_mask=selected_row_mask, - selected_dst_pos=selected_dst_pos, - batch_size=old_batch_size, - HAS_INPUT_IDS=model_input.input_ids is not None, - HAS_B_POSITION_DELTA=model_input.b_position_delta is not None, - HAS_B_SHARED_SEQ_LEN=model_input.b_shared_seq_len is not None, - BLOCK_SIZE=BLOCK_SIZE, - num_warps=8, - num_stages=1, - ) - - model_input.input_ids = out_input_ids - model_input.b_req_idx = out_b_req_idx - model_input.b_mtp_index = out_b_mtp_index - model_input.b_seq_len = out_b_seq_len - model_input.b_position_delta = out_b_position_delta - model_input.b_shared_seq_len = out_b_shared_seq_len - model_input.b_mark_shared_group = _rebuild_mtp_group_markers( - out_b_req_idx, - max_request_rows=max_draft_step + 1, - ) - - if model_input.mtp_draft_input_hiddens is not None: - assert model_input.mtp_draft_input_hiddens.is_cuda - model_input.mtp_draft_input_hiddens = _pack_selected_hidden( - model_input.mtp_draft_input_hiddens, - selected_row_mask, - selected_dst_pos, - dynamic_batch_size, - ) - model_input.batch_size = dynamic_batch_size - - return model_input - - -def prepare_dynamic_mtp_model_input( - model_input: ModelInput, - req_num: int, - dynamic_batch_size: int, - req_to_next_token_scores: torch.Tensor, - pre_draft_step: Optional[int] = None, -): - req_num = int(req_num) - dynamic_batch_size = int(dynamic_batch_size) - assert not model_input.is_prefill, "prepare_dynamic_mtp_model_input only supports decode inputs" - assert req_to_next_token_scores is not None - assert dynamic_batch_size >= req_num - assert dynamic_batch_size <= model_input.batch_size - max_draft_step = MtpManager.get_instance().get_decode_draft_step(is_draft_model=False) - pre_draft_step = max_draft_step if pre_draft_step is None else int(pre_draft_step) - assert 0 <= pre_draft_step <= max_draft_step - assert model_input.batch_size == req_num * (max_draft_step + 1) - assert dynamic_batch_size <= req_num * (pre_draft_step + 1) - - # All compaction work stays on the current CUDA stream and needs no host sync. - model_input.to_cuda() - assert model_input.mem_indexes.shape[0] == dynamic_batch_size - - selected_row_mask = sample_dynamic_mtp_row_mask( - dynamic_batch_size=dynamic_batch_size, - b_req_idx=model_input.b_req_idx, - req_to_next_token_scores=req_to_next_token_scores, - max_draft_step=max_draft_step, - pre_draft_step=pre_draft_step, - ) - - model_input = _compact_decode_model_input( - model_input=model_input, - selected_row_mask=selected_row_mask, - dynamic_batch_size=dynamic_batch_size, - max_draft_step=max_draft_step, - ) - # Decode and draft-cache commit use the compacted b_position_delta, so - # placeholder multimodal metadata only needs to keep ModelInput shapes - # consistent. - if model_input.multimodal_params is not None: - # Read-only placeholders: avoid rebuilding hundreds of nested Python - # objects on every compacted decode iteration. - empty_multimodal_params = {"images": [], "audios": []} - model_input.multimodal_params = [empty_multimodal_params] * dynamic_batch_size - - model_input.max_q_seq_len = 1 - return model_input, selected_row_mask - - @triton.jit def _fwd_kernel_gen_b_req_mtp_start_loc( b_mtp_index, diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 474fc9c9d8..2b3f2302e1 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -76,7 +76,7 @@ def prepare_decode_model_input( if plan.dynamic_batch_size == plan.origin_batch_size: return model_input, None - from lightllm.common.basemodel.triton_kernel.mtp_utils import prepare_dynamic_mtp_model_input + from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import prepare_dynamic_mtp_model_input from lightllm.server.router.model_infer.infer_batch import g_infer_context # mem_indexes 是本轮 decode 新申请、尚未绑定请求和 token 位置的 KV slot。 diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py index 88a2e4b999..5c65d5e5a2 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py @@ -1,12 +1,149 @@ +import json + import torch import pytest import triton import numpy as np +from lightllm.common.basemodel.batch_objs import ModelInput +from lightllm.common.basemodel.mtp_manager import MtpManager +from lightllm.common.basemodel.triton_kernel import dynamic_mtp_utils from lightllm.common.basemodel.triton_kernel.dynamic_mtp_utils import ( _fwd_kernel_cumprod_scores, sample_dynamic_mtp_row_mask, ) +from lightllm.utils.envs_utils import get_env_start_args + + +@pytest.fixture(autouse=True) +def _reset_mtp_manager(): + MtpManager._instance = None + yield + MtpManager._instance = None + + +def test_compact_dynamic_mtp_model_input(monkeypatch): + monkeypatch.setenv( + "LIGHTLLM_START_ARGS", + json.dumps( + { + "diverse_mode": False, + "llm_kv_type": "fp16", + "mtp_mode": "eagle3", + "mtp_dynamic_verify": True, + "mtp_step": 3, + "llm_decode_att_backend": "triton", + } + ), + ) + monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "4") + get_env_start_args.cache_clear() + dynamic_mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() + + model_input = ModelInput( + batch_size=12, + total_token_num=54, + max_q_seq_len=1, + max_kv_seq_len=6, + input_ids=torch.arange(12, dtype=torch.int64, device="cuda") + 1000, + b_req_idx=torch.tensor([0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2], dtype=torch.int32, device="cuda"), + b_mtp_index=torch.tensor([0, 1, 2, 3, 0, 1, 2, 3, 0, 1, 2, 3], dtype=torch.int32, device="cuda"), + b_seq_len=torch.tensor([3, 4, 5, 6, 3, 4, 5, 6, 3, 4, 5, 6], dtype=torch.int32, device="cuda"), + b_position_delta=torch.arange(12, dtype=torch.int32, device="cuda") + 200, + b_shared_seq_len=torch.tensor([0, 0, 0, 0, 7, 7, 7, 7, 9, 9, 9, 9], dtype=torch.int32, device="cuda"), + b_mark_shared_group=torch.tensor([0, 0, 0, 4, 0, 0, 0, 4, 0, 0, 0, 4], dtype=torch.int32, device="cuda"), + mem_indexes=torch.arange(8, dtype=torch.int32, device="cuda") + 100, + mem_indexes_cpu=torch.arange(8, dtype=torch.int32, device="cpu") + 100, + is_prefill=False, + multimodal_params=[{"row": i, "images": [], "audios": []} for i in range(12)], + mtp_draft_input_hiddens=(torch.arange(12 * 5, dtype=torch.float32, device="cuda").reshape(12, 5) + 0.5), + ) + req_to_next_token_scores = torch.tensor( + [ + [1.0, 0.95, 0.90, 0.10, 0.0, 0.0], + [1.0, 0.20, 0.80, 0.80, 0.0, 0.0], + [1.0, 0.99, 0.99, 0.99, 0.0, 0.0], + ], + dtype=torch.float32, + device="cuda", + ) + + compacted_input, selected_row_mask = dynamic_mtp_utils.prepare_dynamic_mtp_model_input( + model_input=model_input, + req_num=3, + dynamic_batch_size=8, + req_to_next_token_scores=req_to_next_token_scores, + ) + torch.cuda.synchronize() + + expected_selected_mask = torch.tensor([1, 1, 1, 0, 1, 0, 0, 0, 1, 1, 1, 1], dtype=torch.int32) + expected_selected_rows = torch.where(expected_selected_mask == 1)[0] + + assert torch.equal(selected_row_mask.cpu(), expected_selected_mask) + assert compacted_input.batch_size == 8 + assert compacted_input.max_q_seq_len == 1 + assert compacted_input.multimodal_params == [{"images": [], "audios": []}] * 8 + + assert torch.equal( + compacted_input.input_ids.cpu(), torch.arange(12, dtype=torch.int64)[expected_selected_rows] + 1000 + ) + assert torch.equal(compacted_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 1, 2, 2, 2, 2], dtype=torch.int32)) + assert torch.equal(compacted_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 0, 0, 1, 2, 3], dtype=torch.int32)) + assert torch.equal(compacted_input.b_seq_len.cpu(), torch.tensor([3, 4, 5, 3, 3, 4, 5, 6], dtype=torch.int32)) + assert torch.equal( + compacted_input.b_position_delta.cpu(), + torch.tensor([200, 201, 202, 204, 208, 209, 210, 211], dtype=torch.int32), + ) + assert torch.equal( + compacted_input.b_shared_seq_len.cpu(), torch.tensor([0, 0, 0, 7, 9, 9, 9, 9], dtype=torch.int32) + ) + assert torch.equal( + compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 1, 0, 0, 0, 4], dtype=torch.int32) + ) + assert torch.equal( + compacted_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 103, 104, 105, 106, 107], dtype=torch.int32) + ) + # CPU/GPU mem indexes are trimmed by SpecEngine.prepare_decode_model_input + # before entering this lower-level row compaction helper. + assert torch.equal(compacted_input.mem_indexes_cpu, torch.arange(8, dtype=torch.int32) + 100) + + expected_hiddens = (torch.arange(12 * 5, dtype=torch.float32).reshape(12, 5) + 0.5)[expected_selected_rows] + assert torch.equal(compacted_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) + + +def test_compaction_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size(monkeypatch): + monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "3") + dynamic_mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() + + model_input = ModelInput( + batch_size=5, + total_token_num=36, + max_q_seq_len=1, + max_kv_seq_len=6, + input_ids=None, + b_req_idx=torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32, device="cuda"), + b_mtp_index=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device="cuda"), + b_seq_len=torch.tensor([3, 4, 5, 6, 7], dtype=torch.int32, device="cuda"), + b_shared_seq_len=None, + b_mark_shared_group=torch.tensor([0, 0, 0, 0, 5], dtype=torch.int32, device="cuda"), + mem_indexes=torch.arange(5, dtype=torch.int32, device="cuda"), + mem_indexes_cpu=torch.arange(5, dtype=torch.int32, device="cpu"), + is_prefill=False, + multimodal_params=[{"images": [], "audios": []} for _ in range(5)], + ) + selected_row_mask = torch.ones((5,), dtype=torch.int32, device="cuda") + + compacted_input = dynamic_mtp_utils._compact_decode_model_input( + model_input=model_input, + selected_row_mask=selected_row_mask, + dynamic_batch_size=5, + max_draft_step=4, + ) + torch.cuda.synchronize() + + assert torch.equal(compacted_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32)) + assert torch.equal(compacted_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32)) + assert torch.equal(compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 0, 2], dtype=torch.int32)) def _reference_cumprod_scores(req_to_next_token_scores, b_req_idx, max_draft_step: int) -> torch.Tensor: diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py index 2e1669570a..c6919b5258 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -1,145 +1,10 @@ -import json import pytest import torch if not torch.cuda.is_available(): pytest.skip("requires CUDA", allow_module_level=True) -from lightllm.common.basemodel.batch_objs import ModelInput -from lightllm.common.basemodel.mtp_manager import MtpManager from lightllm.common.basemodel.triton_kernel import mtp_utils -from lightllm.utils.envs_utils import get_env_start_args - - -@pytest.fixture(autouse=True) -def _reset_mtp_manager(): - MtpManager._instance = None - yield - MtpManager._instance = None - - -def test_compact_dynamic_mtp_model_input(monkeypatch): - monkeypatch.setenv( - "LIGHTLLM_START_ARGS", - json.dumps( - { - "diverse_mode": False, - "llm_kv_type": "fp16", - "mtp_mode": "eagle3", - "mtp_dynamic_verify": True, - "mtp_step": 3, - "llm_decode_att_backend": "triton", - } - ), - ) - monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "4") - get_env_start_args.cache_clear() - mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() - - model_input = ModelInput( - batch_size=12, - total_token_num=54, - max_q_seq_len=1, - max_kv_seq_len=6, - input_ids=torch.arange(12, dtype=torch.int64, device="cuda") + 1000, - b_req_idx=torch.tensor([0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2], dtype=torch.int32, device="cuda"), - b_mtp_index=torch.tensor([0, 1, 2, 3, 0, 1, 2, 3, 0, 1, 2, 3], dtype=torch.int32, device="cuda"), - b_seq_len=torch.tensor([3, 4, 5, 6, 3, 4, 5, 6, 3, 4, 5, 6], dtype=torch.int32, device="cuda"), - b_position_delta=torch.arange(12, dtype=torch.int32, device="cuda") + 200, - b_shared_seq_len=torch.tensor([0, 0, 0, 0, 7, 7, 7, 7, 9, 9, 9, 9], dtype=torch.int32, device="cuda"), - b_mark_shared_group=torch.tensor([0, 0, 0, 4, 0, 0, 0, 4, 0, 0, 0, 4], dtype=torch.int32, device="cuda"), - mem_indexes=torch.arange(8, dtype=torch.int32, device="cuda") + 100, - mem_indexes_cpu=torch.arange(8, dtype=torch.int32, device="cpu") + 100, - is_prefill=False, - multimodal_params=[{"row": i, "images": [], "audios": []} for i in range(12)], - mtp_draft_input_hiddens=(torch.arange(12 * 5, dtype=torch.float32, device="cuda").reshape(12, 5) + 0.5), - ) - req_to_next_token_scores = torch.tensor( - [ - [1.0, 0.95, 0.90, 0.10, 0.0, 0.0], - [1.0, 0.20, 0.80, 0.80, 0.0, 0.0], - [1.0, 0.99, 0.99, 0.99, 0.0, 0.0], - ], - dtype=torch.float32, - device="cuda", - ) - - compacted_input, selected_row_mask = mtp_utils.prepare_dynamic_mtp_model_input( - model_input=model_input, - req_num=3, - dynamic_batch_size=8, - req_to_next_token_scores=req_to_next_token_scores, - ) - torch.cuda.synchronize() - - expected_selected_mask = torch.tensor([1, 1, 1, 0, 1, 0, 0, 0, 1, 1, 1, 1], dtype=torch.int32) - expected_selected_rows = torch.where(expected_selected_mask == 1)[0] - - assert torch.equal(selected_row_mask.cpu(), expected_selected_mask) - assert compacted_input.batch_size == 8 - assert compacted_input.max_q_seq_len == 1 - assert compacted_input.multimodal_params == [{"images": [], "audios": []}] * 8 - - assert torch.equal( - compacted_input.input_ids.cpu(), torch.arange(12, dtype=torch.int64)[expected_selected_rows] + 1000 - ) - assert torch.equal(compacted_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 1, 2, 2, 2, 2], dtype=torch.int32)) - assert torch.equal(compacted_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 0, 0, 1, 2, 3], dtype=torch.int32)) - assert torch.equal(compacted_input.b_seq_len.cpu(), torch.tensor([3, 4, 5, 3, 3, 4, 5, 6], dtype=torch.int32)) - assert torch.equal( - compacted_input.b_position_delta.cpu(), - torch.tensor([200, 201, 202, 204, 208, 209, 210, 211], dtype=torch.int32), - ) - assert torch.equal( - compacted_input.b_shared_seq_len.cpu(), torch.tensor([0, 0, 0, 7, 9, 9, 9, 9], dtype=torch.int32) - ) - assert torch.equal( - compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 1, 0, 0, 0, 4], dtype=torch.int32) - ) - assert torch.equal( - compacted_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 103, 104, 105, 106, 107], dtype=torch.int32) - ) - # CPU/GPU mem indexes are trimmed by SpecEngine.prepare_decode_model_input - # before entering this lower-level row compaction helper. - assert torch.equal(compacted_input.mem_indexes_cpu, torch.arange(8, dtype=torch.int32) + 100) - - expected_hiddens = (torch.arange(12 * 5, dtype=torch.float32).reshape(12, 5) + 0.5)[expected_selected_rows] - assert torch.equal(compacted_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) - - -def test_compaction_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size(monkeypatch): - monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "3") - mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() - - model_input = ModelInput( - batch_size=5, - total_token_num=36, - max_q_seq_len=1, - max_kv_seq_len=6, - input_ids=None, - b_req_idx=torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32, device="cuda"), - b_mtp_index=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device="cuda"), - b_seq_len=torch.tensor([3, 4, 5, 6, 7], dtype=torch.int32, device="cuda"), - b_shared_seq_len=None, - b_mark_shared_group=torch.tensor([0, 0, 0, 0, 5], dtype=torch.int32, device="cuda"), - mem_indexes=torch.arange(5, dtype=torch.int32, device="cuda"), - mem_indexes_cpu=torch.arange(5, dtype=torch.int32, device="cpu"), - is_prefill=False, - multimodal_params=[{"images": [], "audios": []} for _ in range(5)], - ) - selected_row_mask = torch.ones((5,), dtype=torch.int32, device="cuda") - - compacted_input = mtp_utils._compact_decode_model_input( - model_input=model_input, - selected_row_mask=selected_row_mask, - dynamic_batch_size=5, - max_draft_step=4, - ) - torch.cuda.synchronize() - - assert torch.equal(compacted_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32)) - assert torch.equal(compacted_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32)) - assert torch.equal(compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 0, 2], dtype=torch.int32)) def test_mtp_verify_scatter_and_start_locations(): diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 7b7e51f6f6..be2569e466 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -537,7 +537,7 @@ def test_engine_lets_planner_count_requests_with_a_previous_proposal(): def test_dynamic_prepare_keeps_prefix_mem_indexes_and_frees_unused_tail(monkeypatch): - from lightllm.common.basemodel.triton_kernel import mtp_utils as common_mtp_utils + from lightllm.common.basemodel.triton_kernel import dynamic_mtp_utils from lightllm.server.router.model_infer import infer_batch as infer_batch_module from lightllm.server.router.model_infer.mtp_speculative import engine as engine_module @@ -551,7 +551,7 @@ def test_dynamic_prepare_keeps_prefix_mem_indexes_and_frees_unused_tail(monkeypa ) selected_row_mask = torch.tensor([1, 0, 1, 0], dtype=torch.int32) monkeypatch.setattr( - common_mtp_utils, + dynamic_mtp_utils, "prepare_dynamic_mtp_model_input", lambda model_input, **kwargs: (model_input, selected_row_mask), ) From 36261cfba045584c5966fbefdb2784f056733851 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 07:55:33 +0000 Subject: [PATCH 056/103] refactor: clarify MTP verify batch shape --- .../basemodel/triton_kernel/mtp_utils.py | 21 ++++++++++++------- 1 file changed, 13 insertions(+), 8 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index cb13fad7f4..edb4fa15a6 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -16,14 +16,18 @@ def _fwd_kernel_mtp_verify( b_req_mtp_start_loc, b_req_idx, accepted_index, - req_mtp_all_num, + verify_batch_size, BLOCK_SIZE: tl.constexpr, ): cur_index = tl.program_id(0) req_nums = tl.num_programs(axis=0) req_start_loc = tl.load(b_req_mtp_start_loc + cur_index) - req_start_end = tl.load(b_req_mtp_start_loc + cur_index + 1, mask=cur_index + 1 < req_nums, other=req_mtp_all_num) + req_start_end = tl.load( + b_req_mtp_start_loc + cur_index + 1, + mask=cur_index + 1 < req_nums, + other=verify_batch_size, + ) req_mtp_num = req_start_end - req_start_loc cur_req_idx = tl.load(b_req_idx + req_start_loc) @@ -58,20 +62,21 @@ def mtp_verify( Args: req_to_next_token_ids: (max_req_num, verify_width) b_req_mtp_start_loc: (num_reqs,) - new_next_token_ids: (batch_size,) - b_req_idx: (batch_size,) + new_next_token_ids: (verify_batch_size,) + b_req_idx: (verify_batch_size,) Returns: mtp_accept_len: (num_reqs,) - accepted_index: (batch_size,) + accepted_index: (verify_batch_size,) accepted_index: [1, 0, 1, 1, 0], 0 means the token is not accepted, 1 means the token is accepted. """ verify_width = req_to_next_token_ids.shape[1] BLOCK_SIZE = 16 assert verify_width <= BLOCK_SIZE, f"verify_width must be less than {BLOCK_SIZE}" num_reqs = b_req_mtp_start_loc.shape[0] - req_mtp_all_num = b_req_idx.shape[0] + verify_batch_size = b_req_idx.shape[0] + assert new_next_token_ids.shape == b_req_idx.shape mtp_accept_len = torch.empty((num_reqs,), dtype=torch.int32, device=req_to_next_token_ids.device) - accepted_index = torch.empty((req_mtp_all_num,), dtype=torch.int32, device=req_to_next_token_ids.device) + accepted_index = torch.empty((verify_batch_size,), dtype=torch.int32, device=req_to_next_token_ids.device) grid = (num_reqs,) num_warps = 1 @@ -83,7 +88,7 @@ def mtp_verify( b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, accepted_index=accepted_index, - req_mtp_all_num=req_mtp_all_num, + verify_batch_size=verify_batch_size, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps, num_stages=1, From 1dffe62724930b19d56adbf16c001ae5d41109f6 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 08:47:58 +0000 Subject: [PATCH 057/103] refactor: unify MTP proposal token layout --- .../basemodel/triton_kernel/mtp_utils.py | 98 +++++++++++-------- .../mode_backend/chunked_prefill/impl.py | 1 + .../mode_backend/dp_backend/impl.py | 24 +++-- .../dp_overlap_proposers/eagle_no_att.py | 2 + .../dp_overlap_proposers/eagle_utils.py | 49 ++++++---- .../dp_overlap_proposers/eagle_with_att.py | 2 + .../dp_overlap_proposers/vanilla_no_att.py | 12 ++- .../dp_overlap_proposers/vanilla_utils.py | 23 +++-- .../dp_overlap_proposers/vanilla_with_att.py | 12 ++- .../dp_proposers/vanilla_no_att.py | 10 +- .../dp_proposers/vanilla_with_att.py | 10 +- .../mtp_speculative/planner/dspark.py | 9 +- .../mtp_speculative/proposers/base.py | 8 +- .../mtp_speculative/proposers/dflash.py | 15 +-- .../mtp_speculative/proposers/dspark.py | 21 +--- .../mtp_speculative/proposers/eagle_utils.py | 17 ++-- .../proposers/vanilla_no_att.py | 2 + .../proposers/vanilla_utils.py | 14 +-- .../proposers/vanilla_with_att.py | 2 + .../model_infer/mtp_speculative/utils.py | 18 ++-- .../basemodel/triton_kernel/test_mtp_utils.py | 14 +-- .../test_dp_overlap_spec_engine.py | 24 ++--- .../mtp_speculative/test_eagle_overlap.py | 10 +- .../mtp_speculative/test_planner.py | 14 +-- .../mtp_speculative/test_vanilla_overlap.py | 34 +++---- unit_tests/utils/test_speculative_utils.py | 14 ++- 26 files changed, 261 insertions(+), 198 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index edb4fa15a6..f4297866de 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -100,8 +100,9 @@ def mtp_verify( def _fwd_kernel_mtp_scatter_next_token_ids( req_to_next_token_ids, req_to_next_token_ids_stride, - all_next_token_ids, - all_next_token_ids_stride, + target_next_token_ids, + draft_token_ids, + draft_token_ids_stride, req_to_next_token_scores, req_to_next_token_scores_stride, schedule_scores, @@ -109,7 +110,7 @@ def _fwd_kernel_mtp_scatter_next_token_ids( mtp_accept_len, b_req_mtp_start_loc, b_req_idx, - proposal_width, + draft_step, verify_width, HAS_NEXT_TOKEN_SCORES: tl.constexpr, BLOCK_SIZE: tl.constexpr, @@ -122,52 +123,67 @@ def _fwd_kernel_mtp_scatter_next_token_ids( offset = tl.arange(0, BLOCK_SIZE) selected_row = req_start_loc + accept_len - 1 + # 第 0 列写入当前请求最终接受位置对应的 target token。 + target_token_id = tl.load(target_next_token_ids + selected_row) + tl.store( + req_to_next_token_ids + cur_req_idx * req_to_next_token_ids_stride, + target_token_id, + ) + + # 从第 1 列开始写入 draft token;固定宽度 buffer 的未使用列填 1。 + draft_token_id = tl.load( + draft_token_ids + cur_index * draft_token_ids_stride + offset, + mask=offset < draft_step, + other=1, + ) + tl.store( + req_to_next_token_ids + cur_req_idx * req_to_next_token_ids_stride + offset + 1, + draft_token_id, + mask=offset + 1 < verify_width, + ) + if HAS_NEXT_TOKEN_SCORES: - # schedule_scores omits the guaranteed target column. Insert its 1.0 - # here and clear the unused tail of the fixed-width request buffer. - schedule_offset = tl.maximum(offset - 1, 0) + # Target token 必然接受,因此第 0 列分数固定为 1.0。 + tl.store( + req_to_next_token_scores + cur_req_idx * req_to_next_token_scores_stride, + 1.0, + ) + draft_scores = tl.load( - schedule_scores + selected_row * schedule_scores_stride + schedule_offset, - mask=(offset > 0) & (offset < proposal_width), + schedule_scores + cur_index * schedule_scores_stride + offset, + mask=offset < draft_step, other=0.0, ) - next_token_scores = tl.where(offset == 0, 1.0, draft_scores) tl.store( - req_to_next_token_scores + cur_req_idx * req_to_next_token_scores_stride + offset, - next_token_scores, - mask=offset < verify_width, + req_to_next_token_scores + cur_req_idx * req_to_next_token_scores_stride + offset + 1, + draft_scores, + mask=offset + 1 < verify_width, ) - scatter_next_token_ids = tl.load( - all_next_token_ids + selected_row * all_next_token_ids_stride + offset, - mask=offset < proposal_width, - other=1, - ) - tl.store( - req_to_next_token_ids + cur_req_idx * req_to_next_token_ids_stride + offset, - scatter_next_token_ids, - mask=offset < verify_width, - ) return def mtp_scatter_next_token_ids( - req_to_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - all_next_token_ids: torch.Tensor, - b_req_idx: torch.Tensor, - mtp_accept_len: torch.Tensor, - req_to_next_token_scores: Optional[torch.Tensor] = None, - schedule_scores: Optional[torch.Tensor] = None, + req_to_next_token_ids: torch.Tensor, # [max_req_num, verify_width] + b_req_mtp_start_loc: torch.Tensor, # [req_num] + target_next_token_ids: torch.Tensor, # [verify_batch_size] + draft_token_ids: torch.Tensor, # [req_num, draft_step] + b_req_idx: torch.Tensor, # [verify_batch_size] + mtp_accept_len: torch.Tensor, # [req_num] + req_to_next_token_scores: Optional[torch.Tensor] = None, # [max_req_num, verify_width] + schedule_scores: Optional[torch.Tensor] = None, # [req_num, draft_step] ): + """将 target token 和 draft proposal 写入请求级 MTP buffer。""" + verify_width = req_to_next_token_ids.shape[1] BLOCK_SIZE = 16 assert verify_width <= BLOCK_SIZE, f"verify_width must be less than {BLOCK_SIZE}" num_reqs = b_req_mtp_start_loc.shape[0] - proposal_width = all_next_token_ids.shape[1] - assert proposal_width <= verify_width + draft_step = draft_token_ids.shape[1] + assert draft_token_ids.shape[0] == num_reqs + assert draft_step < verify_width if req_to_next_token_scores is not None: assert schedule_scores is not None - assert schedule_scores.shape == (all_next_token_ids.shape[0], proposal_width - 1) + assert schedule_scores.shape == draft_token_ids.shape HAS_NEXT_TOKEN_SCORES = req_to_next_token_scores is not None # Triton launch arguments cannot be None; static verification uses an unused placeholder. @@ -186,8 +202,9 @@ def mtp_scatter_next_token_ids( _fwd_kernel_mtp_scatter_next_token_ids[grid]( req_to_next_token_ids=req_to_next_token_ids, req_to_next_token_ids_stride=req_to_next_token_ids.stride(0), - all_next_token_ids=all_next_token_ids, - all_next_token_ids_stride=all_next_token_ids.stride(0), + target_next_token_ids=target_next_token_ids, + draft_token_ids=draft_token_ids, + draft_token_ids_stride=draft_token_ids.stride(0), req_to_next_token_scores=req_to_next_token_scores_arg, req_to_next_token_scores_stride=req_to_next_token_scores_stride, schedule_scores=schedule_scores_arg, @@ -195,7 +212,7 @@ def mtp_scatter_next_token_ids( mtp_accept_len=mtp_accept_len, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, - proposal_width=proposal_width, + draft_step=draft_step, verify_width=verify_width, HAS_NEXT_TOKEN_SCORES=HAS_NEXT_TOKEN_SCORES, BLOCK_SIZE=BLOCK_SIZE, @@ -314,14 +331,17 @@ def test_mtp_verify(): b_req_idx = torch.tensor([0, 0, 2, 2, 2], dtype=torch.int32, device="cuda") b_req_mtp_start_loc = torch.tensor([0, 2], dtype=torch.int32, device="cuda") new_next_token_ids = torch.tensor([1, 4, 2, 4, 13], dtype=torch.int64, device="cuda") - all_next_token_ids = torch.tensor( - [[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], [13, 14, 15]], dtype=torch.int64, device="cuda" - ) + draft_token_ids = torch.tensor([[2, 3], [8, 9]], dtype=torch.int64, device="cuda") mtp_accept_len, accepted_index = mtp_verify( req_to_next_token_ids, b_req_mtp_start_loc, new_next_token_ids, b_req_idx ) mtp_scatter_next_token_ids( - req_to_next_token_ids, b_req_mtp_start_loc, all_next_token_ids, b_req_idx, mtp_accept_len + req_to_next_token_ids, + b_req_mtp_start_loc, + new_next_token_ids, + draft_token_ids, + b_req_idx, + mtp_accept_len, ) print(mtp_accept_len) print(req_to_next_token_ids) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index f80d86a200..b28d436114 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -314,6 +314,7 @@ def decode_mtp( mtp_utils.scatter_mtp_next_tokens( backend=self, proposal=proposal, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx, mtp_accept_len=mtp_accept_len, diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 7de0541cd4..d21c9ed8ae 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -640,6 +640,9 @@ def _draft_decode_vanilla( mtp_accept_len: torch.Tensor, req_num: int, ): + if b_req_mtp_start_loc is None: + b_req_mtp_start_loc = torch.empty((0,), dtype=torch.int32, device=model_input.b_req_idx.device) + mtp_accept_len = torch.empty_like(b_req_mtp_start_loc) padded_next_token_ids = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input.batch_size, @@ -651,16 +654,17 @@ def _draft_decode_vanilla( main_model_output=model_output, next_token_ids=padded_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=mtp_accept_len, ) if req_num > 0: mtp_utils.scatter_mtp_next_tokens( backend=self, proposal=proposal, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx[:req_num], mtp_accept_len=mtp_accept_len, - valid_row_count=req_num, ) return proposal.extra_mem_indexes_cpu @@ -709,15 +713,18 @@ def _draft_decode_eagle( b_req_mtp_start_loc=padded_start_locs, accept_len=padded_accept_len, ) + proposal.token_ids = proposal.token_ids[:real_request_num] + if getattr(proposal, "schedule_scores", None) is not None: + proposal.schedule_scores = proposal.schedule_scores[:real_request_num] if req_num > 0: mtp_utils.scatter_mtp_next_tokens( backend=self, proposal=proposal, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=model_input.b_req_idx[:req_num], mtp_accept_len=mtp_accept_len, - valid_row_count=req_num, ) return proposal.extra_mem_indexes_cpu @@ -828,7 +835,6 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf _, ) = padded_overlap_prepare_decode_inputs(decode_reqs) req_num0, req_num1 = len(run_reqs0), len(run_reqs1) - all_next_token_ids = [] b_mtp_index_cpu0 = model_input0.b_mtp_index b_mtp_index_cpu1 = model_input1.b_mtp_index with torch.cuda.stream(g_infer_context.get_overlap_stream()): @@ -886,8 +892,6 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf key="mtp_accept_len", gpu_tensor=mtp_accept_len, ) - all_next_token_ids.append(next_token_ids) - verify_event = torch.cuda.Event() verify_event.record() @@ -971,6 +975,10 @@ def _draft_decode_vanilla_overlap( req_num0: int = 0, req_num1: int = 0, ): + if mtp_accept_len is None: + mtp_accept_len = torch.empty((0,), dtype=torch.int32, device=model_input0.b_req_idx.device) + verify_width = self.max_draft_step + 1 + real_request_num0 = req_num0 // verify_width padded_next_token_ids0 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input0.batch_size, @@ -991,18 +999,19 @@ def _draft_decode_vanilla_overlap( main_model_output0=model_output0, next_token_ids0=padded_next_token_ids0, real_verify_rows0=req_num0, - accept_len0=None, + accept_len0=mtp_accept_len[:real_request_num0], main_model_input1=model_input1, main_model_output1=model_output1, next_token_ids1=padded_next_token_ids1, real_verify_rows1=req_num1, - accept_len1=None, + accept_len1=mtp_accept_len[real_request_num0:], ) if req_num0 + req_num1 > 0: mtp_utils.scatter_mtp_next_tokens( backend=self, proposal=proposal, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, @@ -1076,6 +1085,7 @@ def _draft_decode_eagle_overlap( mtp_utils.scatter_mtp_next_tokens( backend=self, proposal=proposal, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 5029b27e8c..19491ab297 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -83,10 +83,12 @@ def propose_next_overlap( main_model_output0, next_token_ids0, real_verify_rows0, + accept_len0, main_model_input1, main_model_output1, next_token_ids1, real_verify_rows1, + accept_len1, draft_step, lambda token_ids: token_ids, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 3724edf8a9..4034de8353 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -113,13 +113,8 @@ def propose_next_dp_eagle_autoregressive_overlap( target_hidden=model_output.mtp_collector.spec_hidden, ) - verify_row_count = real_verify_rows0 + real_verify_rows1 - proposal_token_ids = next_token_ids0.new_full( - (verify_row_count, draft_step + 1), - fill_value=1, - ) - proposal_token_ids[:real_verify_rows0, 0].copy_(next_token_ids0[:real_verify_rows0]) - proposal_token_ids[real_verify_rows0:, 0].copy_(next_token_ids1[:real_verify_rows1]) + total_real_request_count = sum(real_request_counts) + proposal_token_ids = next_token_ids0.new_empty((total_real_request_count, draft_step)) draft_model = proposer.backend.draft_models[0] extend_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) @@ -128,8 +123,7 @@ def propose_next_dp_eagle_autoregressive_overlap( draft_hiddens_by_batch = [] draft_seq_lens_by_batch = [] draft_req_indices_by_batch = [] - proposal_rows_by_batch = [] - proposal_row_offsets = (0, real_verify_rows0) + proposal_row_offsets = (0, real_request_counts[0]) for batch_index, (model_input, extend_output, accepted_tail_rows, real_request_count) in enumerate( zip(model_inputs, extend_outputs, accepted_tail_rows_by_batch, real_request_counts) ): @@ -139,9 +133,10 @@ def propose_next_dp_eagle_autoregressive_overlap( draft_hiddens_by_batch.append(extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows)) draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) - proposal_rows = accepted_tail_rows[:real_request_count] + proposal_row_offsets[batch_index] - proposal_rows_by_batch.append(proposal_rows) - proposal_token_ids[proposal_rows, 1] = draft_token_ids[:real_request_count] + proposal_row_start = proposal_row_offsets[batch_index] + proposal_token_ids[proposal_row_start : proposal_row_start + real_request_count, 0] = draft_token_ids[ + :real_request_count + ] if draft_step == 1: return EagleSpecProposal( @@ -167,7 +162,6 @@ def propose_next_dp_eagle_autoregressive_overlap( empty_multimodal_params = {"images": [], "audios": []} model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size - total_real_request_count = sum(real_request_counts) extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * (draft_step - 1)) extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) hold_mem_index = proposer.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX @@ -197,7 +191,10 @@ def propose_next_dp_eagle_autoregressive_overlap( draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden draft_seq_lens_by_batch[batch_index].add_(1) real_request_count = real_request_counts[batch_index] - proposal_token_ids[proposal_rows_by_batch[batch_index], step + 1] = draft_token_ids[:real_request_count] + proposal_row_start = proposal_row_offsets[batch_index] + proposal_token_ids[proposal_row_start : proposal_row_start + real_request_count, step] = draft_token_ids[ + :real_request_count + ] return EagleSpecProposal( token_ids=proposal_token_ids, @@ -212,14 +209,16 @@ def propose_next_dp_eagle_fixed_layout_overlap( main_model_output0: ModelOutput, next_token_ids0: torch.Tensor, real_verify_rows0: int, + accept_len0: torch.Tensor, main_model_input1: ModelInput, main_model_output1: ModelOutput, next_token_ids1: torch.Tensor, real_verify_rows1: int, + accept_len1: torch.Tensor, draft_step: int, map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], ) -> EagleSpecProposal: - """保持 expanded verify-row layout 运行 DP EAGLE overlap decode。""" + """以 expanded verify-row layout 运行 decode,返回按真实请求压缩的 proposal。""" verify_width = proposer.backend.max_draft_step + 1 model_inputs = (main_model_input0, main_model_input1) @@ -228,9 +227,12 @@ def propose_next_dp_eagle_fixed_layout_overlap( request_capacities_by_batch = tuple(model_input.batch_size // verify_width for model_input in model_inputs) total_real_request_count = sum(real_request_counts) - proposal_token_ids = next_token_ids0.new_empty((sum(real_verify_row_counts), draft_step + 1)) - proposal_token_ids[:real_verify_rows0, 0] = next_token_ids0[:real_verify_rows0] - proposal_token_ids[real_verify_rows0:, 0] = next_token_ids1[:real_verify_rows1] + proposal_token_ids = next_token_ids0.new_empty((total_real_request_count, draft_step)) + proposal_row_offsets = (0, real_request_counts[0]) + accepted_tail_rows_by_batch = ( + torch.arange(0, main_model_input0.batch_size, verify_width, device=next_token_ids0.device) + accept_len0 - 1, + torch.arange(0, main_model_input1.batch_size, verify_width, device=next_token_ids1.device) + accept_len1 - 1, + ) extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * draft_step) extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) @@ -282,8 +284,15 @@ def propose_next_dp_eagle_fixed_layout_overlap( ) draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden - proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids_by_batch[0][:real_verify_rows0] - proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids_by_batch[1][:real_verify_rows1] + for batch_index, draft_token_ids in enumerate(draft_token_ids_by_batch): + real_request_count = real_request_counts[batch_index] + proposal_row_start = proposal_row_offsets[batch_index] + proposal_token_ids[ + proposal_row_start : proposal_row_start + real_request_count, step + ] = draft_token_ids.index_select( + 0, + accepted_tail_rows_by_batch[batch_index][:real_request_count].long(), + ) return EagleSpecProposal( token_ids=proposal_token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 8e15e1ee1d..2cb95f7ac5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -83,10 +83,12 @@ def propose_next_overlap( main_model_output0, next_token_ids0, real_verify_rows0, + accept_len0, main_model_input1, main_model_output1, next_token_ids1, real_verify_rows1, + accept_len1, draft_step, lambda token_ids: token_ids, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 7ddf7e22ee..08f3436803 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -52,7 +52,15 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) + return propose_next_chained_mtp( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + ) def propose_next_overlap( self, @@ -74,9 +82,11 @@ def propose_next_overlap( main_model_output0, next_token_ids0, real_verify_rows0, + accept_len0, main_model_input1, main_model_output1, next_token_ids1, real_verify_rows1, + accept_len1, draft_step, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py index 5a663dab29..bdbab7f1ce 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py @@ -48,24 +48,31 @@ def propose_next_dp_chained_mtp_overlap( main_model_output0: ModelOutput, next_token_ids0: torch.Tensor, real_verify_rows0: int, + accept_len0: torch.Tensor, main_model_input1: ModelInput, main_model_output1: ModelOutput, next_token_ids1: torch.Tensor, real_verify_rows1: int, + accept_len1: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: - """为两个 DP microbatch 运行 Vanilla chained overlap decode。""" + """为两个 DP microbatch 运行 decode,返回按真实请求压缩的 proposal。""" model_inputs = (main_model_input0, main_model_input1) + verify_width = proposer.backend.max_draft_step + 1 real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) + real_request_counts = tuple(row_count // verify_width for row_count in real_verify_rows) + accepted_tail_rows = ( + torch.arange(0, real_verify_rows0, verify_width, device=next_token_ids0.device) + accept_len0 - 1, + torch.arange(0, real_verify_rows1, verify_width, device=next_token_ids1.device) + accept_len1 - 1, + ) draft_token_ids = [next_token_ids0, next_token_ids1] draft_hiddens = [ main_model_output0.mtp_collector.spec_hidden, main_model_output1.mtp_collector.spec_hidden, ] - proposal_token_ids = next_token_ids0.new_empty((sum(real_verify_rows), draft_step + 1)) - proposal_token_ids[:real_verify_rows0, 0] = next_token_ids0[:real_verify_rows0] - proposal_token_ids[real_verify_rows0:, 0] = next_token_ids1[:real_verify_rows1] + proposal_token_ids = next_token_ids0.new_empty((sum(real_request_counts), draft_step)) + request_offset = real_request_counts[0] for step in range(draft_step): for batch_index, model_input in enumerate(model_inputs): @@ -77,8 +84,12 @@ def propose_next_dp_chained_mtp_overlap( draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden draft_token_ids[batch_index] = proposer.backend._gen_argmax_token_ids(draft_output) - proposal_token_ids[:real_verify_rows0, step + 1] = draft_token_ids[0][:real_verify_rows0] - proposal_token_ids[real_verify_rows0:, step + 1] = draft_token_ids[1][:real_verify_rows1] + proposal_token_ids[:request_offset, step] = draft_token_ids[0].index_select( + 0, accepted_tail_rows[0][:request_offset].long() + ) + proposal_token_ids[request_offset:, step] = draft_token_ids[1].index_select( + 0, accepted_tail_rows[1][: real_request_counts[1]].long() + ) return VanillaSpecProposal( token_ids=proposal_token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 6050d544fa..751b44bc16 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -52,7 +52,15 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) + return propose_next_chained_mtp( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + ) def propose_next_overlap( self, @@ -74,9 +82,11 @@ def propose_next_overlap( main_model_output0, next_token_ids0, real_verify_rows0, + accept_len0, main_model_input1, main_model_output1, next_token_ids1, real_verify_rows1, + accept_len1, draft_step, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py index 2ec72367d0..7a4b4f2779 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -29,4 +29,12 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) + return propose_next_chained_mtp( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py index 9b1bf82be5..3782bd339a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -29,4 +29,12 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp(self, main_model_input, main_model_output, next_token_ids, draft_step) + return propose_next_chained_mtp( + self, + main_model_input, + main_model_output, + next_token_ids, + b_req_mtp_start_loc, + draft_step, + accept_len, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index a04134c2e3..552644bb83 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -103,13 +103,8 @@ def _update_confidence_probs(self, confidence_probs, req_num: int) -> None: if draft_confidence_probs.size == 0: return - # Confidence is scattered onto one accepted-tail row per request; - # unused verify rows remain zero. - valid_rows = np.any(draft_confidence_probs > 0.0, axis=1) - if not np.any(valid_rows): - return - - conditional_probs = np.clip(draft_confidence_probs[valid_rows], 0.01, 0.99) + # Proposal scores are dense request rows and contain draft columns only. + conditional_probs = np.clip(draft_confidence_probs, 0.01, 0.99) survival_scores = np.cumprod(conditional_probs, axis=1) dynamic_batch_size = self._select_dynamic_batch_size_from_survival_scores( req_num=int(req_num), diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 065149bfb5..3d45c5a1d9 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -34,8 +34,9 @@ def __post_init__(self) -> None: class SpecProposal: """Common candidate-token output produced by every MTP proposer. - `token_ids` has shape `[verify_batch, draft_step + 1]`; column 0 contains - target-model tokens and the remaining columns contain draft candidates. + `token_ids` has shape `[req_num, draft_step]` and contains only draft-model + candidates. The target model's latest token is kept separately and merged + into the request-level MTP buffer only when the proposal is persisted. `extra_mem_indexes_cpu` uniformly tracks every KV slot considered for release, including rejected target rows and proposal-owned temporary rows. Mode-specific scheduling metadata belongs to the corresponding subclass. @@ -96,7 +97,8 @@ def propose_next( by dynamic scheduling. `b_req_mtp_start_loc` identifies each logical request's first row. - Column 0 of the returned proposal must equal `next_token_ids`. + The returned proposal contains one dense row per logical request and + only the `draft_step` candidate tokens produced by the draft model. """ raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index ecfb218080..c907608da3 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -55,14 +55,8 @@ def propose_next( accept_len: torch.Tensor | None = None, ) -> DFlashSpecProposal: request_count = int(b_req_mtp_start_loc.shape[0]) - verify_row_count = int(next_token_ids.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - proposal_token_ids = next_token_ids.new_full( - (verify_row_count, draft_step + 1), - fill_value=1, - ) - proposal_token_ids[:, 0] = next_token_ids # One accepted-tail anchor expands to a complete block-diffusion draft. accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() @@ -85,17 +79,12 @@ def propose_next( else: flat_draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) - proposal_token_ids[accepted_tail_rows, 1:] = block_draft_token_ids[:, :draft_step] + proposal_token_ids = block_draft_token_ids[:, :draft_step].contiguous() schedule_scores = None if self.enable_dynmaic_mtp: block_draft_token_probs = flat_draft_token_probs.reshape(request_count, block_size) - schedule_scores = torch.zeros( - (verify_row_count, draft_step), - dtype=torch.float32, - device=next_token_ids.device, - ) - schedule_scores[accepted_tail_rows] = block_draft_token_probs[:, :draft_step].float() + schedule_scores = block_draft_token_probs[:, :draft_step].float().contiguous() return DFlashSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 02134ede99..1d35cf29c5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -58,23 +58,9 @@ def propose_next( accept_len: torch.Tensor | None = None, ) -> DSparkSpecProposal: request_count = int(b_req_mtp_start_loc.shape[0]) - verify_row_count = int(next_token_ids.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - proposal_token_ids = next_token_ids.new_full( - (verify_row_count, draft_step + 1), - fill_value=1, - ) - proposal_token_ids[:, 0] = next_token_ids - schedule_scores = ( - torch.zeros( - (verify_row_count, draft_step), - dtype=torch.float32, - device=next_token_ids.device, - ) - if self.enable_dynmaic_mtp - else None - ) + schedule_scores = None accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() draft_input, extra_mem_indexes_cpu = build_parallel_block_draft_input( @@ -96,7 +82,7 @@ def propose_next( else: flat_draft_token_ids = draft_output.mtp_collector.draft_token_ids block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) - proposal_token_ids[accepted_tail_rows, 1:] = block_draft_token_ids[:, :draft_step] + proposal_token_ids = block_draft_token_ids[:, :draft_step].contiguous() if self.enable_dynmaic_mtp: confidence_logits = draft_output.mtp_collector.confidence_logits @@ -104,13 +90,14 @@ def propose_next( raise RuntimeError("DSpark dynamic verify requires confidence head logits") # Match the clamp used by the GPU dynamic row selector before it # converts conditional confidence to prefix survival probability. - schedule_scores[accepted_tail_rows] = ( + schedule_scores = ( confidence_logits[:, :draft_step] .sigmoid() .clamp( min=0.01, max=0.99, ) + .contiguous() ) schedule_scores_cpu = None diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 561e544bea..ab940ab81b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -90,17 +90,12 @@ def propose_next_eagle( ) -> EagleSpecProposal: """运行 EAGLE extend 后接单 token decode 的通用 proposal 流程。""" - verify_row_count = int(next_token_ids.shape[0]) request_count = int(b_req_mtp_start_loc.shape[0]) - proposal_token_ids = next_token_ids.new_full( - (verify_row_count, draft_step + 1), - fill_value=1, - ) - proposal_token_ids[:, 0].copy_(next_token_ids) + proposal_token_ids = next_token_ids.new_empty((request_count, draft_step)) collect_schedule_scores = proposer.enable_dynmaic_mtp schedule_scores = ( torch.zeros( - (verify_row_count, draft_step), + (request_count, draft_step), dtype=torch.float32, device=next_token_ids.device, ) @@ -131,14 +126,14 @@ def propose_next_eagle( model_output=accepted_tail_output, map_draft_token_ids=map_draft_token_ids, ) - schedule_scores[accepted_tail_rows, 0] = draft_token_probs.float() + schedule_scores[:, 0] = draft_token_probs.float() else: draft_token_ids = generate_eagle_token_ids( proposer=proposer, model_output=accepted_tail_output, map_draft_token_ids=map_draft_token_ids, ) - proposal_token_ids[accepted_tail_rows, 1] = draft_token_ids + proposal_token_ids[:, 0] = draft_token_ids draft_hidden = extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows) if draft_step == 1: @@ -181,14 +176,14 @@ def propose_next_eagle( model_output=draft_output, map_draft_token_ids=map_draft_token_ids, ) - schedule_scores[accepted_tail_rows, step] = draft_token_probs.float() + schedule_scores[:, step] = draft_token_probs.float() else: draft_token_ids = generate_eagle_token_ids( proposer=proposer, model_output=draft_output, map_draft_token_ids=map_draft_token_ids, ) - proposal_token_ids[accepted_tail_rows, step + 1] = draft_token_ids + proposal_token_ids[:, step] = draft_token_ids draft_hidden = draft_output.mtp_collector.spec_hidden draft_seq_lens.add_(1) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index 13a1efa67e..7934535c9d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -39,5 +39,7 @@ def propose_next( main_model_input=main_model_input, main_model_output=main_model_output, next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, + accept_len=accept_len, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py index 0d4953d9cd..a640928350 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py @@ -45,18 +45,20 @@ def propose_next_chained_mtp( main_model_input: ModelInput, main_model_output: ModelOutput, next_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, draft_step: int, + accept_len: torch.Tensor, ) -> VanillaSpecProposal: """依次运行 Vanilla chained MTP 模块并生成 proposal。""" - verify_row_count = int(next_token_ids.shape[0]) + request_count = int(b_req_mtp_start_loc.shape[0]) + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() draft_token_ids = next_token_ids draft_hidden = main_model_output.mtp_collector.spec_hidden - proposal_token_ids = next_token_ids.new_empty((verify_row_count, draft_step + 1)) - proposal_token_ids[:, 0] = next_token_ids + proposal_token_ids = next_token_ids.new_empty((request_count, draft_step)) schedule_scores = ( torch.empty( - (verify_row_count, draft_step), + (request_count, draft_step), dtype=torch.float32, device=next_token_ids.device, ) @@ -72,10 +74,10 @@ def propose_next_chained_mtp( draft_hidden = draft_output.mtp_collector.spec_hidden if proposer.enable_dynmaic_mtp: draft_token_ids, draft_token_probs = proposer.backend._gen_argmax_token_ids_and_prob(draft_output) - schedule_scores[:, step] = draft_token_probs + schedule_scores[:, step] = draft_token_probs.index_select(0, accepted_tail_rows) else: draft_token_ids = proposer.backend._gen_argmax_token_ids(draft_output) - proposal_token_ids[:, step + 1] = draft_token_ids + proposal_token_ids[:, step] = draft_token_ids.index_select(0, accepted_tail_rows) return VanillaSpecProposal( token_ids=proposal_token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index 23fc7eaeb1..cd0fa579bb 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -39,5 +39,7 @@ def propose_next( main_model_input=main_model_input, main_model_output=main_model_output, next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, + accept_len=accept_len, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index c6cfa48858..571a22d551 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -64,11 +64,11 @@ def verify_mtp_tokens( def scatter_mtp_next_tokens( backend: ModeBackend, - proposal: SpecProposal, - b_req_mtp_start_loc: torch.Tensor, - b_req_idx: torch.Tensor, - mtp_accept_len: torch.Tensor, - valid_row_count: Optional[int] = None, + proposal: SpecProposal, # proposal.token_ids: [req_num, draft_step] + target_next_token_ids: torch.Tensor, # [verify_batch_size] + b_req_mtp_start_loc: torch.Tensor, # [req_num] + b_req_idx: torch.Tensor, # [verify_batch_size] + mtp_accept_len: torch.Tensor, # [req_num] ) -> None: """Persist the next MTP proposal and optional scheduling scores by request.""" @@ -77,19 +77,15 @@ def scatter_mtp_next_tokens( from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import EagleSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal - token_ids = proposal.token_ids scored_proposal_types = (VanillaSpecProposal, EagleSpecProposal, DFlashSpecProposal, DSparkSpecProposal) schedule_scores = proposal.schedule_scores if isinstance(proposal, scored_proposal_types) else None - if valid_row_count is not None: - token_ids = token_ids[:valid_row_count] - if schedule_scores is not None: - schedule_scores = schedule_scores[:valid_row_count] sampling_params_manager = backend.model.req_manager.req_sampling_params_manager mtp_scatter_next_token_ids( req_to_next_token_ids=sampling_params_manager.req_to_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=token_ids, + target_next_token_ids=target_next_token_ids, + draft_token_ids=proposal.token_ids, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, req_to_next_token_scores=( diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py index c6919b5258..b8206c78ad 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -17,14 +17,14 @@ def test_mtp_verify_scatter_and_start_locations(): b_mtp_index = torch.tensor([0, 1, 0, 1, 2], dtype=torch.int32, device="cuda") b_req_mtp_start_loc = mtp_utils.gen_b_req_mtp_start_loc(b_mtp_index, num_reqs=2) new_next_token_ids = torch.tensor([1, 4, 2, 4, 13], dtype=torch.int64, device="cuda") - all_next_token_ids = torch.tensor( - [[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], [13, 14, 15]], + draft_token_ids = torch.tensor( + [[2, 3], [8, 9]], dtype=torch.int64, device="cuda", ) req_to_next_token_scores = torch.full((3, 5), -1.0, dtype=torch.float32, device="cuda") schedule_scores = torch.tensor( - [[0.9, 0.8], [0.8, 0.7], [0.7, 0.6], [0.6, 0.5], [0.5, 0.4]], + [[0.9, 0.8], [0.7, 0.6]], dtype=torch.float32, device="cuda", ) @@ -35,7 +35,8 @@ def test_mtp_verify_scatter_and_start_locations(): mtp_utils.mtp_scatter_next_token_ids( req_to_next_token_ids=req_to_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - all_next_token_ids=all_next_token_ids, + target_next_token_ids=new_next_token_ids, + draft_token_ids=draft_token_ids, b_req_idx=b_req_idx, mtp_accept_len=mtp_accept_len, req_to_next_token_scores=req_to_next_token_scores, @@ -49,7 +50,7 @@ def test_mtp_verify_scatter_and_start_locations(): assert torch.equal( req_to_next_token_ids.cpu(), torch.tensor( - [[1, 2, 3, 1, 1], [1, 2, 0, -1, -1], [7, 8, 9, 1, 1]], + [[1, 2, 3, 1, 1], [1, 2, 0, -1, -1], [2, 8, 9, 1, 1]], dtype=torch.int64, ), ) @@ -69,7 +70,8 @@ def test_mtp_scatter_handles_zero_draft_step(): mtp_utils.mtp_scatter_next_token_ids( req_to_next_token_ids=req_to_next_token_ids, b_req_mtp_start_loc=torch.tensor([0], dtype=torch.int32, device="cuda"), - all_next_token_ids=torch.tensor([[42]], dtype=torch.int64, device="cuda"), + target_next_token_ids=torch.tensor([42], dtype=torch.int64, device="cuda"), + draft_token_ids=torch.empty((1, 0), dtype=torch.int64, device="cuda"), b_req_idx=torch.tensor([0], dtype=torch.int32, device="cuda"), mtp_accept_len=torch.tensor([1], dtype=torch.int32, device="cuda"), req_to_next_token_scores=req_to_next_token_scores, diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index b7c724fbf2..dbabd15bdc 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -22,7 +22,7 @@ def __init__(self): def propose_next(self, **kwargs): self.propose_args = kwargs - token_ids = kwargs["next_token_ids"].new_zeros((16, 8)) + token_ids = kwargs["next_token_ids"].new_zeros((kwargs["b_req_mtp_start_loc"].shape[0], 7)) return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([123], dtype=torch.int32))], @@ -30,8 +30,8 @@ def propose_next(self, **kwargs): def propose_next_overlap(self, **kwargs): self.propose_overlap_args = kwargs - row_count = kwargs["real_verify_rows0"] + kwargs["real_verify_rows1"] - token_ids = kwargs["next_token_ids0"].new_zeros((row_count, 8)) + request_count = (kwargs["real_verify_rows0"] + kwargs["real_verify_rows1"]) // 8 + token_ids = kwargs["next_token_ids0"].new_zeros((request_count, 7)) return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([456], dtype=torch.int32))], @@ -133,8 +133,8 @@ def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(monkeypatch): assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) assert torch.equal(propose_args["b_req_mtp_start_loc"], torch.tensor([0, 8], dtype=torch.int32)) assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) - assert scatter_args["proposal"].token_ids.shape == (16, 8) - assert scatter_args["valid_row_count"] == 8 + assert scatter_args["proposal"].token_ids.shape == (1, 7) + assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([123], dtype=torch.int32)) @@ -162,8 +162,8 @@ def test_dp_vanilla_uses_dp_engine_proposer(monkeypatch): propose_args = backend.spec_engine.propose_args assert propose_args["next_token_ids"].shape == (16,) assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) - assert scatter_args["proposal"].token_ids.shape == (16, 8) - assert scatter_args["valid_row_count"] == 8 + assert scatter_args["proposal"].token_ids.shape == (8, 7) + assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([123], dtype=torch.int32)) @@ -209,7 +209,8 @@ def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(monkeypat assert propose_args["real_verify_rows1"] == 16 assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) - assert scatter_args["proposal"].token_ids.shape == (24, 8) + assert scatter_args["proposal"].token_ids.shape == (3, 7) + assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([456], dtype=torch.int32)) @@ -248,7 +249,8 @@ def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(monkeypatch) assert propose_args["next_token_ids1"].shape == (16,) assert propose_args["real_verify_rows0"] == 8 assert propose_args["real_verify_rows1"] == 16 - assert propose_args["accept_len0"] is None - assert propose_args["accept_len1"] is None - assert scatter_args["proposal"].token_ids.shape == (24, 8) + assert torch.equal(propose_args["accept_len0"], torch.tensor([2], dtype=torch.int32)) + assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) + assert scatter_args["proposal"].token_ids.shape == (3, 7) + assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([456], dtype=torch.int32)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index fd717ca820..22a9752e79 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -104,11 +104,10 @@ def test_overlap_eagle_keeps_fixed_verify_layout(monkeypatch): assert draft_model.extend_batch_sizes is None assert draft_model.decode_batch_sizes == [(6, 6), (6, 6)] - assert proposal.token_ids.shape == (9, 3) - assert torch.equal(proposal.token_ids[:, 0], torch.tensor([0, 1, 2, 10, 11, 12, 13, 14, 15])) - expected_draft_tokens = torch.tensor([0, 1, 2, 0, 1, 2, 3, 4, 5]) + assert proposal.token_ids.shape == (3, 2) + expected_draft_tokens = torch.tensor([1, 0, 5]) + assert torch.equal(proposal.token_ids[:, 0], expected_draft_tokens) assert torch.equal(proposal.token_ids[:, 1], expected_draft_tokens) - assert torch.equal(proposal.token_ids[:, 2], expected_draft_tokens) assert len(proposal.extra_mem_indexes_cpu) == 1 assert torch.equal(proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, torch.arange(6, dtype=torch.int32)) assert proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None @@ -162,7 +161,8 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): assert draft_model.decode_inputs[0][1] is model_input1 assert draft_model.extend_batch_sizes == (6, 6) assert draft_model.decode_batch_sizes == [(2, 2)] - assert proposal.token_ids.shape == (9, 3) + assert proposal.token_ids.shape == (3, 2) + assert torch.equal(proposal.token_ids, torch.tensor([[1, 0], [0, 0], [5, 1]])) assert len(proposal.extra_mem_indexes_cpu) == 1 assert torch.equal(proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, torch.arange(3, dtype=torch.int32)) assert proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index be2569e466..c7d2c900ae 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -177,24 +177,26 @@ def test_scatter_mtp_next_tokens_consumes_mode_proposal(monkeypatch): ) ) proposal = DFlashSpecProposal( - token_ids=torch.arange(6, dtype=torch.int64).view(3, 2), + token_ids=torch.arange(2, dtype=torch.int64).view(2, 1), extra_mem_indexes_cpu=[], - schedule_scores=torch.arange(3, dtype=torch.float32).view(3, 1), + schedule_scores=torch.arange(2, dtype=torch.float32).view(2, 1), ) + next_token_ids = torch.tensor([10, 11], dtype=torch.int64) mtp_utils.scatter_mtp_next_tokens( backend=backend, proposal=proposal, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=torch.tensor([0, 1], dtype=torch.int32), b_req_idx=torch.tensor([0, 1], dtype=torch.int32), mtp_accept_len=torch.ones(2, dtype=torch.int32), - valid_row_count=2, ) assert scatter_args["req_to_next_token_ids"] is req_to_next_token_ids assert scatter_args["req_to_next_token_scores"] is req_to_next_token_scores - assert torch.equal(scatter_args["all_next_token_ids"], proposal.token_ids[:2]) - assert torch.equal(scatter_args["schedule_scores"], proposal.schedule_scores[:2]) + assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) + assert torch.equal(scatter_args["draft_token_ids"], proposal.token_ids) + assert torch.equal(scatter_args["schedule_scores"], proposal.schedule_scores) def test_dp_planner_returns_fixed_backend_draft_step(): @@ -365,7 +367,7 @@ def test_eagle_proposer_skips_draft_forward_for_zero_steps(): ) assert isinstance(proposal, EagleSpecProposal) - assert proposal.token_ids.tolist() == [[10], [11]] + assert proposal.token_ids.shape == (2, 0) assert proposal.schedule_scores.shape == (2, 0) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py index 8cb4761989..0cb075b80f 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -26,40 +26,38 @@ def microbatch_overlap_decode(self, input0, input1): def test_dp_vanilla_proposer_owns_overlap_decode(): draft_models = [_DraftModel(), _DraftModel()] backend = SimpleNamespace( + max_draft_step=2, draft_models=draft_models, _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) - model_input0 = SimpleNamespace(batch_size=4) - model_input1 = SimpleNamespace(batch_size=4) + model_input0 = SimpleNamespace(batch_size=6) + model_input1 = SimpleNamespace(batch_size=6) proposal = proposer.propose_next_overlap( main_model_input0=model_input0, main_model_output0=ModelOutput( - logits=torch.empty((4, 1)), - mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((4, 2))), + logits=torch.empty((6, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), ), - next_token_ids0=torch.tensor([10, 11, 0, 0], dtype=torch.int64), - real_verify_rows0=2, - accept_len0=None, + next_token_ids0=torch.tensor([10, 11, 0, 0, 0, 0], dtype=torch.int64), + real_verify_rows0=3, + accept_len0=torch.tensor([2], dtype=torch.int32), main_model_input1=model_input1, main_model_output1=ModelOutput( - logits=torch.empty((4, 1)), - mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((4, 2))), + logits=torch.empty((6, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), ), - next_token_ids1=torch.tensor([20, 21, 22, 0], dtype=torch.int64), + next_token_ids1=torch.tensor([20, 21, 22, 0, 0, 0], dtype=torch.int64), real_verify_rows1=3, - accept_len1=None, + accept_len1=torch.tensor([1], dtype=torch.int32), draft_step=2, ) assert proposal.token_ids.tolist() == [ - [10, 0, 0], - [11, 1, 1], - [20, 0, 0], - [21, 1, 1], - [22, 2, 2], + [1, 1], + [0, 0], ] assert proposal.extra_mem_indexes_cpu == [] - assert draft_models[0].decode_batch_sizes == [(4, 4)] - assert draft_models[1].decode_batch_sizes == [(4, 4)] + assert draft_models[0].decode_batch_sizes == [(6, 6)] + assert draft_models[1].decode_batch_sizes == [(6, 6)] diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 0b57b2e86d..b480c13db8 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -228,7 +228,6 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(monkeypatch) block_size = 4 max_draft_step = 3 verify_row_count = 5 - accepted_tail_rows = torch.tensor([0, 3]) flat_draft_token_ids = torch.arange(2 * block_size) flat_draft_token_probs = torch.arange(2 * block_size, dtype=torch.float32) / 10 @@ -262,13 +261,12 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(monkeypatch) assert isinstance(proposal, DFlashSpecProposal) expected_blocks = flat_draft_token_ids.reshape(2, block_size)[:, :2] - torch.testing.assert_close(proposal.token_ids[accepted_tail_rows, 1:], expected_blocks) - assert proposal.schedule_scores.shape == (verify_row_count, 2) - for step in range(2): - scores = proposal.schedule_scores[:, step] - expected_scores = torch.zeros(verify_row_count) - expected_scores[accepted_tail_rows] = flat_draft_token_probs.reshape(2, block_size)[:, step] - torch.testing.assert_close(scores, expected_scores) + torch.testing.assert_close(proposal.token_ids, expected_blocks) + assert proposal.schedule_scores.shape == (2, 2) + torch.testing.assert_close( + proposal.schedule_scores, + flat_draft_token_probs.reshape(2, block_size)[:, :2], + ) def test_dflash_reuses_decode_input_for_kv_commit(): From 14e7bbae5ccb3601f33beeb1a277f0931d5091ef Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 09:06:35 +0000 Subject: [PATCH 058/103] refactor: clarify MTP metric and memory handling --- .../model_infer/mode_backend/chunked_prefill/impl.py | 5 ++--- .../router/model_infer/mode_backend/dp_backend/impl.py | 10 ++++------ .../server/router/model_infer/mtp_speculative/utils.py | 6 +++--- .../router/model_infer/mtp_speculative/test_planner.py | 4 ++-- 4 files changed, 11 insertions(+), 14 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index b28d436114..715f178128 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -363,7 +363,7 @@ def decode_mtp( backend=self, decode_reqs=decode_reqs, accept_lengths_cpu=mtp_accept_len_cpu, - verified_row_reqs=run_reqs, + verify_run_reqs=run_reqs, ) select_mask = accepted_index_cpu.to(dtype=torch.bool) @@ -376,8 +376,7 @@ def decode_mtp( extra_post_req_handle_func=self.extra_post_req_handle_func, ) - proposal.extra_mem_indexes_cpu.insert( - 0, + proposal.extra_mem_indexes_cpu.append( MtpMemIndexesToFree( mem_indexes_cpu=model_input.mem_indexes_cpu, free_mask_cpu=accepted_index_cpu == 0, diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index d21c9ed8ae..23dc585d94 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -593,7 +593,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): backend=self, decode_reqs=decode_reqs, accept_lengths_cpu=mtp_accept_len_cpu, - verified_row_reqs=run_reqs, + verify_run_reqs=run_reqs, ) verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) @@ -601,8 +601,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): # 第三阶段 event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() - extra_mem_indexes_cpu.insert( - 0, + extra_mem_indexes_cpu.append( MtpMemIndexesToFree( mem_indexes_cpu=model_input.mem_indexes_cpu[0:req_num], free_mask_cpu=accepted_index_cpu == 0, @@ -924,7 +923,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf backend=self, decode_reqs=decode_reqs, accept_lengths_cpu=mtp_accept_len_cpu, - verified_row_reqs=run_reqs, + verify_run_reqs=run_reqs, ) verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) @@ -934,8 +933,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf mem_indexes_cpu = torch.cat( (model_input0.mem_indexes_cpu[0:req_num0], model_input1.mem_indexes_cpu[0:req_num1]), dim=0 ) - extra_mem_indexes_cpu.insert( - 0, + extra_mem_indexes_cpu.append( MtpMemIndexesToFree( mem_indexes_cpu=mem_indexes_cpu, free_mask_cpu=accepted_index_cpu == 0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index 571a22d551..7933266918 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -99,7 +99,7 @@ def record_request_mtp_metrics( backend: ModeBackend, decode_reqs: List[InferReq], accept_lengths_cpu: torch.Tensor, - verified_row_reqs: List[InferReq], + verify_run_reqs: List[InferReq], ) -> None: """Accumulate user-visible MTP metrics on each request.""" @@ -108,10 +108,10 @@ def record_request_mtp_metrics( accept_lengths = accept_lengths_cpu.tolist() assert len(accept_lengths) == len(decode_reqs) - verify_rows_by_req = Counter(req.req_idx for req in verified_row_reqs) + verify_count_by_req_idx = Counter(req.req_idx for req in verify_run_reqs) for req, accept_len in zip(decode_reqs, accept_lengths): req.update_mtp_accepted_token_num(accept_token_num=accept_len - 1) - verify_token_num = verify_rows_by_req[req.req_idx] + verify_token_num = verify_count_by_req_idx[req.req_idx] if verify_token_num > 0: req.update_mtp_verify_token_num(verify_token_num=verify_token_num) req.update_mtp_verify_step_num(verify_step_num=1) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index c7d2c900ae..1e5105593c 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -420,7 +420,7 @@ def update_mtp_verify_step_num(self, verify_step_num: int): backend=backend, decode_reqs=[req0, req1], accept_lengths_cpu=torch.tensor([2, 1], dtype=torch.int32), - verified_row_reqs=[req0, req0, req1], + verify_run_reqs=[req0, req0, req1], ) assert (req0.accepted, req0.verified, req0.verify_steps) == (1, 2, 1) @@ -431,7 +431,7 @@ def update_mtp_verify_step_num(self, verify_step_num: int): backend=backend, decode_reqs=[fixed_req], accept_lengths_cpu=torch.tensor([3], dtype=torch.int32), - verified_row_reqs=[fixed_req] * 4, + verify_run_reqs=[fixed_req] * 4, ) assert (fixed_req.accepted, fixed_req.verified, fixed_req.verify_steps) == (2, 4, 1) From 79e61968bbf49c4664f858a75835d21b2f380479 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 09:14:15 +0000 Subject: [PATCH 059/103] refactor: rename draft KV state initialization --- .../model_infer/mode_backend/chunked_prefill/impl.py | 2 +- .../model_infer/mode_backend/dp_backend/impl.py | 4 ++-- .../router/model_infer/mtp_speculative/dp_engine.py | 4 ++-- .../model_infer/mtp_speculative/dp_overlap_engine.py | 4 ++-- .../mtp_speculative/dp_overlap_proposers/base.py | 2 +- .../mtp_speculative/dp_overlap_proposers/eagle3.py | 12 ++++++------ .../dp_overlap_proposers/eagle_no_att.py | 12 ++++++------ .../dp_overlap_proposers/eagle_utils.py | 2 +- .../dp_overlap_proposers/eagle_with_att.py | 12 ++++++------ .../dp_overlap_proposers/vanilla_no_att.py | 12 ++++++------ .../dp_overlap_proposers/vanilla_utils.py | 2 +- .../dp_overlap_proposers/vanilla_with_att.py | 12 ++++++------ .../mtp_speculative/dp_proposers/eagle3.py | 6 +++--- .../mtp_speculative/dp_proposers/eagle_no_att.py | 6 +++--- .../mtp_speculative/dp_proposers/eagle_with_att.py | 6 +++--- .../mtp_speculative/dp_proposers/vanilla_no_att.py | 6 +++--- .../mtp_speculative/dp_proposers/vanilla_with_att.py | 6 +++--- .../router/model_infer/mtp_speculative/engine.py | 4 ++-- .../model_infer/mtp_speculative/proposers/base.py | 2 +- .../model_infer/mtp_speculative/proposers/dflash.py | 6 +++--- .../model_infer/mtp_speculative/proposers/dspark.py | 6 +++--- .../model_infer/mtp_speculative/proposers/eagle3.py | 6 +++--- .../mtp_speculative/proposers/eagle_no_att.py | 6 +++--- .../mtp_speculative/proposers/eagle_utils.py | 2 +- .../mtp_speculative/proposers/eagle_with_att.py | 6 +++--- .../proposers/parallel_block_utils.py | 2 +- .../mtp_speculative/proposers/vanilla_no_att.py | 6 +++--- .../mtp_speculative/proposers/vanilla_utils.py | 2 +- .../mtp_speculative/proposers/vanilla_with_att.py | 6 +++--- .../model_infer/mtp_speculative/test_planner.py | 2 +- 30 files changed, 83 insertions(+), 83 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 715f178128..3b49508907 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -209,7 +209,7 @@ def prefill_mtp( ) # mtp kv fill spec_engine = self.spec_engine - spec_engine.build_draft_state_from_prefill( + spec_engine.fill_draft_model_kv_state( target_model_input=model_input, target_model_output=model_output, next_token_ids=next_token_ids, diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 23dc585d94..6462a5eadf 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -478,7 +478,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] copy_len=req_num, device=model_input.b_req_idx.device, ) - self.prefill_draft_engine.build_draft_state_from_prefill( + self.prefill_draft_engine.fill_draft_model_kv_state( target_model_input=model_input, target_model_output=model_output, next_token_ids=draft_next_token_ids_gpu, @@ -785,7 +785,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I source_start=req_num0, ) - self.prefill_draft_engine.build_draft_state_from_prefill_overlap( + self.prefill_draft_engine.fill_draft_model_kv_state_overlap( target_model_input0=model_input0, target_model_output0=model_output0, next_token_ids0=draft_next_token_ids_gpu0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py index 57692cd55e..65d9af6f0d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py @@ -25,13 +25,13 @@ def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: boo ) self.planner: BaseDpPlanner = build_dp_planner(backend=backend) - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - self.proposer.build_draft_state_from_prefill( + self.proposer.fill_draft_model_kv_state( target_model_input=target_model_input, target_model_output=target_model_output, next_token_ids=next_token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index 3fbe894b0b..5f5bf23376 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -28,7 +28,7 @@ def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: boo ) self.planner: BaseDpOverlapPlanner = build_dp_overlap_planner(backend=backend) - def build_draft_state_from_prefill_overlap( + def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -37,7 +37,7 @@ def build_draft_state_from_prefill_overlap( target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - self.proposer.build_draft_state_from_prefill_overlap( + self.proposer.fill_draft_model_kv_state_overlap( target_model_input0=target_model_input0, target_model_output0=target_model_output0, next_token_ids0=next_token_ids0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py index fe25b40b9d..b8ecf3393a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -10,7 +10,7 @@ class BaseDpOverlapProposer(BaseSpecProposer, ABC): """DP proposer 的完整接口,扩展双 microbatch overlap 操作。""" @abstractmethod - def build_draft_state_from_prefill_overlap( + def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index b022c4baa8..b916831a3d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -3,12 +3,12 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( - build_dp_eagle_draft_state_from_prefill_overlap, + fill_dp_eagle_draft_model_kv_state_overlap, propose_next_dp_eagle_autoregressive_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, propose_next_eagle, ) @@ -19,15 +19,15 @@ class DpOverlapEagle3Proposer(BaseDpOverlapProposer): def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) - def build_draft_state_from_prefill_overlap( + def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -36,7 +36,7 @@ def build_draft_state_from_prefill_overlap( target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - build_dp_eagle_draft_state_from_prefill_overlap( + fill_dp_eagle_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 19491ab297..ea3850bc42 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -3,12 +3,12 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( - build_dp_eagle_draft_state_from_prefill_overlap, + fill_dp_eagle_draft_model_kv_state_overlap, propose_next_dp_eagle_fixed_layout_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, propose_next_eagle, ) @@ -16,15 +16,15 @@ class DpOverlapEagleNoAttProposer(BaseDpOverlapProposer): """DP ``eagle_no_att`` proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) - def build_draft_state_from_prefill_overlap( + def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -33,7 +33,7 @@ def build_draft_state_from_prefill_overlap( target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - build_dp_eagle_draft_state_from_prefill_overlap( + fill_dp_eagle_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 4034de8353..420df582a8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -17,7 +17,7 @@ ) -def build_dp_eagle_draft_state_from_prefill_overlap( +def fill_dp_eagle_draft_model_kv_state_overlap( proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_output0: ModelOutput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 2cb95f7ac5..8169c76cc2 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -3,12 +3,12 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( - build_dp_eagle_draft_state_from_prefill_overlap, + fill_dp_eagle_draft_model_kv_state_overlap, propose_next_dp_eagle_fixed_layout_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, propose_next_eagle, ) @@ -16,15 +16,15 @@ class DpOverlapEagleWithAttProposer(BaseDpOverlapProposer): """DP ``eagle_with_att`` proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) - def build_draft_state_from_prefill_overlap( + def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -33,7 +33,7 @@ def build_draft_state_from_prefill_overlap( target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - build_dp_eagle_draft_state_from_prefill_overlap( + fill_dp_eagle_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 08f3436803..aa60533fab 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -3,12 +3,12 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( - build_dp_chained_mtp_draft_state_from_prefill_overlap, + fill_dp_chained_mtp_draft_model_kv_state_overlap, propose_next_dp_chained_mtp_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - build_chained_mtp_draft_state_from_prefill, + fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -16,15 +16,15 @@ class DpOverlapVanillaNoAttProposer(BaseDpOverlapProposer): """DP ``vanilla_no_att`` proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_chained_mtp_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) - def build_draft_state_from_prefill_overlap( + def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -33,7 +33,7 @@ def build_draft_state_from_prefill_overlap( target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - build_dp_chained_mtp_draft_state_from_prefill_overlap( + fill_dp_chained_mtp_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py index bdbab7f1ce..e73834ace8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py @@ -9,7 +9,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal -def build_dp_chained_mtp_draft_state_from_prefill_overlap( +def fill_dp_chained_mtp_draft_model_kv_state_overlap( proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_output0: ModelOutput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 751b44bc16..98044778a5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -3,12 +3,12 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( - build_dp_chained_mtp_draft_state_from_prefill_overlap, + fill_dp_chained_mtp_draft_model_kv_state_overlap, propose_next_dp_chained_mtp_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - build_chained_mtp_draft_state_from_prefill, + fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -16,15 +16,15 @@ class DpOverlapVanillaWithAttProposer(BaseDpOverlapProposer): """DP ``vanilla_with_att`` proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_chained_mtp_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) - def build_draft_state_from_prefill_overlap( + def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, @@ -33,7 +33,7 @@ def build_draft_state_from_prefill_overlap( target_model_output1: ModelOutput, next_token_ids1: torch.Tensor, ) -> None: - build_dp_chained_mtp_draft_state_from_prefill_overlap( + fill_dp_chained_mtp_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py index d4289e0060..c4b68e5b9b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, propose_next_eagle, ) @@ -15,13 +15,13 @@ class DpEagle3Proposer(BaseDpProposer): def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py index 637eef849d..497ed79804 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, propose_next_eagle, ) @@ -12,13 +12,13 @@ class DpEagleNoAttProposer(BaseDpProposer): """普通 DP ``eagle_no_att`` proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py index 0bdfa38be6..0e7fa084b4 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, propose_next_eagle, ) @@ -12,13 +12,13 @@ class DpEagleWithAttProposer(BaseDpProposer): """普通 DP ``eagle_with_att`` proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py index 7a4b4f2779..efbda49a5c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - build_chained_mtp_draft_state_from_prefill, + fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -12,13 +12,13 @@ class DpVanillaNoAttProposer(BaseDpProposer): """普通 DP ``vanilla_no_att`` proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_chained_mtp_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py index 3782bd339a..bfccf0caae 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - build_chained_mtp_draft_state_from_prefill, + fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -12,13 +12,13 @@ class DpVanillaWithAttProposer(BaseDpProposer): """普通 DP ``vanilla_with_att`` proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_chained_mtp_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 2b3f2302e1..f31c4d0812 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -41,13 +41,13 @@ def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: boo # Prefill draft-state initialization. - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - self.proposer.build_draft_state_from_prefill( + self.proposer.fill_draft_model_kv_state( target_model_input=target_model_input, target_model_output=target_model_output, next_token_ids=next_token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 3d45c5a1d9..7017d262ed 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -60,7 +60,7 @@ def __init__(self, *, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> None: self.enable_dynmaic_mtp = bool(enable_dynmaic_mtp) @abstractmethod - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index c907608da3..161d9d95fc 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -12,7 +12,7 @@ ) from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( build_parallel_block_draft_input, - build_parallel_block_draft_state_from_prefill, + fill_parallel_block_draft_model_kv_state, extend_parallel_block_draft_kv_cache, ) @@ -31,13 +31,13 @@ class DFlashProposer(BaseSpecProposer): the accepted-tail anchor and mask-token positions. """ - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_parallel_block_draft_state_from_prefill( + fill_parallel_block_draft_model_kv_state( proposer=self, target_model_input=target_model_input, target_model_output=target_model_output, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 1d35cf29c5..09414edeb6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -13,7 +13,7 @@ ) from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( build_parallel_block_draft_input, - build_parallel_block_draft_state_from_prefill, + fill_parallel_block_draft_model_kv_state, extend_parallel_block_draft_kv_cache, ) @@ -34,13 +34,13 @@ class DSparkProposer(BaseSpecProposer): The confidence head supplies per-position scheduling scores. """ - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_parallel_block_draft_state_from_prefill( + fill_parallel_block_draft_model_kv_state( proposer=self, target_model_input=target_model_input, target_model_output=target_model_output, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py index 2c618bc98d..dc4a07793c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, generate_eagle_token_ids, generate_eagle_token_ids_and_prob, propose_next_eagle, @@ -23,13 +23,13 @@ def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): return generate_eagle_token_ids_and_prob(self, model_output, self._map_draft_token_ids) - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py index bf1896e83e..125e48a568 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, propose_next_eagle, ) @@ -12,13 +12,13 @@ class EagleNoAttProposer(BaseSpecProposer): """不使用 attention KV cache 的 EAGLE proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index ab940ab81b..3461cbda52 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -24,7 +24,7 @@ class EagleSpecProposal(SpecProposal): schedule_scores: torch.Tensor | None = None -def build_eagle_draft_state_from_prefill( +def fill_eagle_draft_model_kv_state( proposer: BaseSpecProposer, target_model_input: ModelInput, target_model_output: ModelOutput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index ab0cb63f5a..42675d93a6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, - build_eagle_draft_state_from_prefill, + fill_eagle_draft_model_kv_state, propose_next_eagle, ) @@ -12,13 +12,13 @@ class EagleWithAttProposer(BaseSpecProposer): """使用 attention KV cache 的 EAGLE proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_eagle_draft_state_from_prefill(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py index d74f0ff28b..3033721c5a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py @@ -12,7 +12,7 @@ @torch.no_grad() -def build_parallel_block_draft_state_from_prefill( +def fill_parallel_block_draft_model_kv_state( proposer: BaseSpecProposer, target_model_input: ModelInput, target_model_output: ModelOutput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index 7934535c9d..790913cc09 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - build_chained_mtp_draft_state_from_prefill, + fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -12,13 +12,13 @@ class VanillaNoAttProposer(BaseSpecProposer): """不使用 attention KV cache 的 Vanilla chained MTP proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_chained_mtp_draft_state_from_prefill( + fill_chained_mtp_draft_model_kv_state( proposer=self, target_model_input=target_model_input, target_model_output=target_model_output, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py index a640928350..56ce760faa 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py @@ -17,7 +17,7 @@ class VanillaSpecProposal(SpecProposal): schedule_scores: torch.Tensor | None = None -def build_chained_mtp_draft_state_from_prefill( +def fill_chained_mtp_draft_model_kv_state( proposer: BaseSpecProposer, target_model_input: ModelInput, target_model_output: ModelOutput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index cd0fa579bb..4efb7eb230 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - build_chained_mtp_draft_state_from_prefill, + fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -12,13 +12,13 @@ class VanillaWithAttProposer(BaseSpecProposer): """使用 attention KV cache 的 Vanilla chained MTP proposer。""" - def build_draft_state_from_prefill( + def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, next_token_ids: torch.Tensor, ) -> None: - build_chained_mtp_draft_state_from_prefill( + fill_chained_mtp_draft_model_kv_state( proposer=self, target_model_input=target_model_input, target_model_output=target_model_output, diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 1e5105593c..7458cd6d2c 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -137,7 +137,7 @@ def test_spec_engine_only_exposes_planning_and_proposal_interfaces(): } assert public_methods == { - "build_draft_state_from_prefill", + "fill_draft_model_kv_state", "plan_decode", "prepare_decode_model_input", "propose_next", From b868bf74bcd529b60a85e8751ce42dc77b59caf0 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 09:22:13 +0000 Subject: [PATCH 060/103] refactor: clarify target token inputs for draft KV fill --- .../model_infer/mode_backend/chunked_prefill/impl.py | 2 +- .../model_infer/mode_backend/dp_backend/impl.py | 12 ++++++------ .../router/model_infer/mtp_speculative/dp_engine.py | 4 ++-- .../model_infer/mtp_speculative/dp_overlap_engine.py | 8 ++++---- .../mtp_speculative/dp_overlap_proposers/base.py | 4 ++-- .../mtp_speculative/dp_overlap_proposers/eagle3.py | 12 ++++++------ .../dp_overlap_proposers/eagle_no_att.py | 12 ++++++------ .../dp_overlap_proposers/eagle_utils.py | 8 ++++---- .../dp_overlap_proposers/eagle_with_att.py | 12 ++++++------ .../dp_overlap_proposers/vanilla_no_att.py | 12 ++++++------ .../dp_overlap_proposers/vanilla_utils.py | 6 +++--- .../dp_overlap_proposers/vanilla_with_att.py | 12 ++++++------ .../mtp_speculative/dp_proposers/eagle3.py | 4 ++-- .../mtp_speculative/dp_proposers/eagle_no_att.py | 4 ++-- .../mtp_speculative/dp_proposers/eagle_with_att.py | 4 ++-- .../mtp_speculative/dp_proposers/vanilla_no_att.py | 4 ++-- .../mtp_speculative/dp_proposers/vanilla_with_att.py | 4 ++-- .../router/model_infer/mtp_speculative/engine.py | 4 ++-- .../model_infer/mtp_speculative/proposers/base.py | 4 ++-- .../model_infer/mtp_speculative/proposers/dflash.py | 4 ++-- .../model_infer/mtp_speculative/proposers/dspark.py | 4 ++-- .../model_infer/mtp_speculative/proposers/eagle3.py | 4 ++-- .../mtp_speculative/proposers/eagle_no_att.py | 4 ++-- .../mtp_speculative/proposers/eagle_utils.py | 4 ++-- .../mtp_speculative/proposers/eagle_with_att.py | 4 ++-- .../proposers/parallel_block_utils.py | 2 +- .../mtp_speculative/proposers/vanilla_no_att.py | 4 ++-- .../mtp_speculative/proposers/vanilla_utils.py | 4 ++-- .../mtp_speculative/proposers/vanilla_with_att.py | 4 ++-- 29 files changed, 85 insertions(+), 85 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 3b49508907..90d477a75e 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -212,7 +212,7 @@ def prefill_mtp( spec_engine.fill_draft_model_kv_state( target_model_input=model_input, target_model_output=model_output, - next_token_ids=next_token_ids, + target_next_token_ids=next_token_ids, ) g_infer_context.copy_linear_att_state_to_cache_buffer( b_req_idx=model_input.b_req_idx, diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 6462a5eadf..dee3d61668 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -472,7 +472,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] ) # mtp kv fill - draft_next_token_ids_gpu = self._build_padded_next_token_ids( + target_next_token_ids_gpu = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input.batch_size, copy_len=req_num, @@ -481,7 +481,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] self.prefill_draft_engine.fill_draft_model_kv_state( target_model_input=model_input, target_model_output=model_output, - next_token_ids=draft_next_token_ids_gpu, + target_next_token_ids=target_next_token_ids_gpu, ) if req_num > 0: g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=b_req_idx, reqs=run_reqs) @@ -770,14 +770,14 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I b_prefill_has_output_cpu=b_has_out_cpu, ) - draft_next_token_ids_gpu0 = self._build_padded_next_token_ids( + target_next_token_ids_gpu0 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input0.batch_size, copy_len=req_num0, device=model_input0.b_req_idx.device, source_start=0, ) - draft_next_token_ids_gpu1 = self._build_padded_next_token_ids( + target_next_token_ids_gpu1 = self._build_padded_next_token_ids( token_ids=next_token_ids, batch_size=model_input1.batch_size, copy_len=req_num1, @@ -788,10 +788,10 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I self.prefill_draft_engine.fill_draft_model_kv_state_overlap( target_model_input0=model_input0, target_model_output0=model_output0, - next_token_ids0=draft_next_token_ids_gpu0, + target_next_token_ids0=target_next_token_ids_gpu0, target_model_input1=model_input1, target_model_output1=model_output1, - next_token_ids1=draft_next_token_ids_gpu1, + target_next_token_ids1=target_next_token_ids_gpu1, ) if req_num0 + req_num1 > 0 and g_infer_context.is_linear_att_mixed_model: diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py index 65d9af6f0d..b2f32003d0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py @@ -29,12 +29,12 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: self.proposer.fill_draft_model_kv_state( target_model_input=target_model_input, target_model_output=target_model_output, - next_token_ids=next_token_ids, + target_next_token_ids=target_next_token_ids, ) def propose_next( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index 5f5bf23376..cb31226dc7 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -32,18 +32,18 @@ def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: self.proposer.fill_draft_model_kv_state_overlap( target_model_input0=target_model_input0, target_model_output0=target_model_output0, - next_token_ids0=next_token_ids0, + target_next_token_ids0=target_next_token_ids0, target_model_input1=target_model_input1, target_model_output1=target_model_output1, - next_token_ids1=next_token_ids1, + target_next_token_ids1=target_next_token_ids1, ) def propose_next_overlap( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py index b8ecf3393a..2b417dc37c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -14,10 +14,10 @@ def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: """Build draft state from two overlapped target-prefill microbatches.""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index b916831a3d..28ad442178 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -23,27 +23,27 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: fill_dp_eagle_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, - next_token_ids0, + target_next_token_ids0, target_model_input1, target_model_output1, - next_token_ids1, + target_next_token_ids1, ) def propose_next( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index ea3850bc42..00039d69e1 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -20,27 +20,27 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: fill_dp_eagle_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, - next_token_ids0, + target_next_token_ids0, target_model_input1, target_model_output1, - next_token_ids1, + target_next_token_ids1, ) def propose_next( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 420df582a8..07ce70666b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -21,10 +21,10 @@ def fill_dp_eagle_draft_model_kv_state_overlap( proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: """使用两个 target prefill microbatch 初始化 EAGLE draft state。""" @@ -32,12 +32,12 @@ def fill_dp_eagle_draft_model_kv_state_overlap( prepare_mtp_prefill_inputs( model_input=target_model_input0, - b_next_token_ids=next_token_ids0, + b_next_token_ids=target_next_token_ids0, mtp_draft_input_hiddens=target_model_output0.mtp_collector.spec_hidden, ) prepare_mtp_prefill_inputs( model_input=target_model_input1, - b_next_token_ids=next_token_ids1, + b_next_token_ids=target_next_token_ids1, mtp_draft_input_hiddens=target_model_output1.mtp_collector.spec_hidden, ) proposer.backend.draft_models[0].microbatch_overlap_prefill(target_model_input0, target_model_input1) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 8169c76cc2..61745960db 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -20,27 +20,27 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: fill_dp_eagle_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, - next_token_ids0, + target_next_token_ids0, target_model_input1, target_model_output1, - next_token_ids1, + target_next_token_ids1, ) def propose_next( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index aa60533fab..4ecb5590c3 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -20,27 +20,27 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: fill_dp_chained_mtp_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, - next_token_ids0, + target_next_token_ids0, target_model_input1, target_model_output1, - next_token_ids1, + target_next_token_ids1, ) def propose_next( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py index e73834ace8..cdd368bb6b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py @@ -13,10 +13,10 @@ def fill_dp_chained_mtp_draft_model_kv_state_overlap( proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: """为两个 DP microbatch 依次构建 Vanilla chained draft state。""" @@ -27,7 +27,7 @@ def fill_dp_chained_mtp_draft_model_kv_state_overlap( target_model_output0.mtp_collector.spec_hidden, target_model_output1.mtp_collector.spec_hidden, ] - draft_token_ids = [next_token_ids0, next_token_ids1] + draft_token_ids = [target_next_token_ids0, target_next_token_ids1] for draft_model in proposer.backend.draft_models: for batch_index, model_input in enumerate(model_inputs): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 98044778a5..301e2088d9 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -20,27 +20,27 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_next_token_ids1: torch.Tensor, ) -> None: fill_dp_chained_mtp_draft_model_kv_state_overlap( self, target_model_input0, target_model_output0, - next_token_ids0, + target_next_token_ids0, target_model_input1, target_model_output1, - next_token_ids1, + target_next_token_ids1, ) def propose_next( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py index c4b68e5b9b..7187051b2a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py @@ -19,9 +19,9 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py index 497ed79804..02c22af68a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py @@ -16,9 +16,9 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py index 0e7fa084b4..3f94fb961a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py @@ -16,9 +16,9 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py index efbda49a5c..930977003e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -16,9 +16,9 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py index bfccf0caae..559aaaa956 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -16,9 +16,9 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index f31c4d0812..c340d6c5cf 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -45,12 +45,12 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: self.proposer.fill_draft_model_kv_state( target_model_input=target_model_input, target_model_output=target_model_output, - next_token_ids=next_token_ids, + target_next_token_ids=target_next_token_ids, ) # Decode planning. diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 7017d262ed..1a144f9351 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -64,7 +64,7 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: """Build draft KV/state from target prefill before the first decode verify. @@ -73,7 +73,7 @@ def fill_draft_model_kv_state( mem_indexes are reused by the draft state builder. - `target_model_output`: target output containing the features needed by the selected speculative algorithm. - - `next_token_ids`: first accepted target token, shape [run_req_num]. + - `target_next_token_ids`: first accepted target token, shape [run_req_num]. This hook only prepares draft-side state. It does not create proposal tokens; the first decode iteration creates them through `propose_next`. diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index 161d9d95fc..e27838dde4 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -35,13 +35,13 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: fill_parallel_block_draft_model_kv_state( proposer=self, target_model_input=target_model_input, target_model_output=target_model_output, - next_token_ids=next_token_ids, + target_next_token_ids=target_next_token_ids, ) @torch.no_grad() diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 09414edeb6..ff3fb0da39 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -38,13 +38,13 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: fill_parallel_block_draft_model_kv_state( proposer=self, target_model_input=target_model_input, target_model_output=target_model_output, - next_token_ids=next_token_ids, + target_next_token_ids=target_next_token_ids, ) @torch.no_grad() diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py index dc4a07793c..a95775859e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py @@ -27,9 +27,9 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py index 125e48a568..c1ca306b87 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -16,9 +16,9 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 3461cbda52..976905e41c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -28,7 +28,7 @@ def fill_eagle_draft_model_kv_state( proposer: BaseSpecProposer, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: """使用 target prefill 输出初始化 EAGLE draft state。""" @@ -36,7 +36,7 @@ def fill_eagle_draft_model_kv_state( prepare_mtp_prefill_inputs( model_input=target_model_input, - b_next_token_ids=next_token_ids, + b_next_token_ids=target_next_token_ids, mtp_draft_input_hiddens=target_model_output.mtp_collector.spec_hidden, ) proposer.backend.draft_models[0].forward(target_model_input) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index 42675d93a6..7cd6ffc701 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -16,9 +16,9 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, next_token_ids) + fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py index 3033721c5a..2a4a5e784e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py @@ -16,7 +16,7 @@ def fill_parallel_block_draft_model_kv_state( proposer: BaseSpecProposer, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: """使用 target hidden 初始化 parallel-block drafter 的 KV state。""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index 790913cc09..58021a295e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -16,13 +16,13 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: fill_chained_mtp_draft_model_kv_state( proposer=self, target_model_input=target_model_input, target_model_output=target_model_output, - next_token_ids=next_token_ids, + target_next_token_ids=target_next_token_ids, ) def propose_next( diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py index 56ce760faa..834f978b39 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py @@ -21,14 +21,14 @@ def fill_chained_mtp_draft_model_kv_state( proposer: BaseSpecProposer, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: """构建 Vanilla chained MTP 各级 draft model 的 prefill state。""" from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_token_ids = next_token_ids + draft_token_ids = target_next_token_ids for draft_model in proposer.backend.draft_models: prepare_mtp_prefill_inputs( model_input=target_model_input, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index 4efb7eb230..559dd0ef8d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -16,13 +16,13 @@ def fill_draft_model_kv_state( self, target_model_input: ModelInput, target_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_next_token_ids: torch.Tensor, ) -> None: fill_chained_mtp_draft_model_kv_state( proposer=self, target_model_input=target_model_input, target_model_output=target_model_output, - next_token_ids=next_token_ids, + target_next_token_ids=target_next_token_ids, ) def propose_next( From 4696ba97752079b015ca43ff071de185bd635385 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Wed, 19 Aug 2026 09:45:13 +0000 Subject: [PATCH 061/103] refactor: clarify target proposal inputs --- .../mode_backend/chunked_prefill/impl.py | 6 +- .../mode_backend/dp_backend/impl.py | 36 ++++++------ .../model_infer/mtp_speculative/dp_engine.py | 16 +++--- .../mtp_speculative/dp_overlap_engine.py | 28 +++++----- .../dp_overlap_proposers/base.py | 16 +++--- .../dp_overlap_proposers/eagle3.py | 36 ++++++------ .../dp_overlap_proposers/eagle_no_att.py | 36 ++++++------ .../dp_overlap_proposers/eagle_utils.py | 56 ++++++++++--------- .../dp_overlap_proposers/eagle_with_att.py | 36 ++++++------ .../dp_overlap_proposers/vanilla_no_att.py | 36 ++++++------ .../dp_overlap_proposers/vanilla_utils.py | 26 ++++----- .../dp_overlap_proposers/vanilla_with_att.py | 36 ++++++------ .../mtp_speculative/dp_proposers/eagle3.py | 12 ++-- .../dp_proposers/eagle_no_att.py | 12 ++-- .../dp_proposers/eagle_with_att.py | 12 ++-- .../dp_proposers/vanilla_no_att.py | 12 ++-- .../dp_proposers/vanilla_with_att.py | 12 ++-- .../model_infer/mtp_speculative/engine.py | 16 +++--- .../mtp_speculative/proposers/base.py | 12 ++-- .../mtp_speculative/proposers/dflash.py | 14 ++--- .../mtp_speculative/proposers/dspark.py | 14 ++--- .../mtp_speculative/proposers/eagle3.py | 12 ++-- .../mtp_speculative/proposers/eagle_no_att.py | 12 ++-- .../mtp_speculative/proposers/eagle_utils.py | 30 +++++----- .../proposers/eagle_with_att.py | 12 ++-- .../proposers/parallel_block_utils.py | 44 ++++++++------- .../proposers/vanilla_no_att.py | 12 ++-- .../proposers/vanilla_utils.py | 20 +++---- .../proposers/vanilla_with_att.py | 12 ++-- .../test_dp_overlap_spec_engine.py | 20 +++---- .../mtp_speculative/test_eagle_overlap.py | 24 ++++---- .../mtp_speculative/test_planner.py | 6 +- .../mtp_speculative/test_vanilla_overlap.py | 12 ++-- unit_tests/utils/test_speculative_utils.py | 10 ++-- 34 files changed, 356 insertions(+), 350 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 90d477a75e..4ed1519e72 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py @@ -304,9 +304,9 @@ def decode_mtp( ) proposal = spec_engine.propose_next( - main_model_input=model_input, - main_model_output=model_output, - next_token_ids=next_token_ids, + target_model_input=model_input, + target_model_output=model_output, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=spec_plan.draft_step, accept_len=mtp_accept_len, diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index dee3d61668..49be5fcced 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -649,9 +649,9 @@ def _draft_decode_vanilla( device=model_input.b_req_idx.device, ) proposal = self.decode_draft_engine.propose_next( - main_model_input=model_input, - main_model_output=model_output, - next_token_ids=padded_next_token_ids, + target_model_input=model_input, + target_model_output=model_output, + target_next_token_ids=padded_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, accept_len=mtp_accept_len, ) @@ -706,9 +706,9 @@ def _draft_decode_eagle( # full-row extend, followed by autoregressive drafting over one row per # (real or HOLD) request. proposal = self.decode_draft_engine.propose_next( - main_model_input=model_input, - main_model_output=model_output, - next_token_ids=padded_next_token_ids, + target_model_input=model_input, + target_model_output=model_output, + target_next_token_ids=padded_next_token_ids, b_req_mtp_start_loc=padded_start_locs, accept_len=padded_accept_len, ) @@ -993,14 +993,14 @@ def _draft_decode_vanilla_overlap( ) proposal = self.decode_draft_engine.propose_next_overlap( - main_model_input0=model_input0, - main_model_output0=model_output0, - next_token_ids0=padded_next_token_ids0, + target_model_input0=model_input0, + target_model_output0=model_output0, + target_next_token_ids0=padded_next_token_ids0, real_verify_rows0=req_num0, accept_len0=mtp_accept_len[:real_request_num0], - main_model_input1=model_input1, - main_model_output1=model_output1, - next_token_ids1=padded_next_token_ids1, + target_model_input1=model_input1, + target_model_output1=model_output1, + target_next_token_ids1=padded_next_token_ids1, real_verify_rows1=req_num1, accept_len1=mtp_accept_len[real_request_num0:], ) @@ -1067,14 +1067,14 @@ def _draft_decode_eagle_overlap( ) proposal = self.decode_draft_engine.propose_next_overlap( - main_model_input0=model_input0, - main_model_output0=model_output0, - next_token_ids0=padded_next_token_ids0, + target_model_input0=model_input0, + target_model_output0=model_output0, + target_next_token_ids0=padded_next_token_ids0, real_verify_rows0=req_num0, accept_len0=padded_accept_len0, - main_model_input1=model_input1, - main_model_output1=model_output1, - next_token_ids1=padded_next_token_ids1, + target_model_input1=model_input1, + target_model_output1=model_output1, + target_next_token_ids1=padded_next_token_ids1, real_verify_rows1=req_num1, accept_len1=padded_accept_len1, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py index b2f32003d0..6b9a109c3f 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py @@ -39,16 +39,16 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - accept_len: Optional[torch.Tensor] = None, + target_model_input: ModelInput, # batch_size = verify_batch_size + target_model_output: ModelOutput, # logits: [verify_batch_size, vocab_size] + target_next_token_ids: torch.Tensor, # [verify_batch_size] + b_req_mtp_start_loc: torch.Tensor, # [req_num] + accept_len: Optional[torch.Tensor] = None, # [req_num] ) -> SpecProposal: return self.proposer.propose_next( - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=self.planner.get_draft_step(), accept_len=accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index cb31226dc7..3c7ad489fe 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -48,26 +48,26 @@ def fill_draft_model_kv_state_overlap( def propose_next_overlap( self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, # batch_size = padded_verify_batch_size0 + target_model_output0: ModelOutput, # logits: [padded_verify_batch_size0, vocab_size] + target_next_token_ids0: torch.Tensor, # [padded_verify_batch_size0] real_verify_rows0: int, - accept_len0: Optional[torch.Tensor], - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + accept_len0: Optional[torch.Tensor], # [req_num0] + target_model_input1: ModelInput, # batch_size = padded_verify_batch_size1 + target_model_output1: ModelOutput, # logits: [padded_verify_batch_size1, vocab_size] + target_next_token_ids1: torch.Tensor, # [padded_verify_batch_size1] real_verify_rows1: int, - accept_len1: Optional[torch.Tensor], + accept_len1: Optional[torch.Tensor], # [req_num1] ) -> SpecProposal: return self.proposer.propose_next_overlap( - main_model_input0=main_model_input0, - main_model_output0=main_model_output0, - next_token_ids0=next_token_ids0, + target_model_input0=target_model_input0, + target_model_output0=target_model_output0, + target_next_token_ids0=target_next_token_ids0, real_verify_rows0=real_verify_rows0, accept_len0=accept_len0, - main_model_input1=main_model_input1, - main_model_output1=main_model_output1, - next_token_ids1=next_token_ids1, + target_model_input1=target_model_input1, + target_model_output1=target_model_output1, + target_next_token_ids1=target_next_token_ids1, real_verify_rows1=real_verify_rows1, accept_len1=accept_len1, draft_step=self.planner.get_draft_step(), diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py index 2b417dc37c..b085dc85bd 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -26,16 +26,16 @@ def fill_draft_model_kv_state_overlap( @abstractmethod def propose_next_overlap( self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, # batch_size = padded_verify_batch_size0 + target_model_output0: ModelOutput, # logits: [padded_verify_batch_size0, vocab_size] + target_next_token_ids0: torch.Tensor, # [padded_verify_batch_size0] real_verify_rows0: int, - accept_len0: torch.Tensor | None, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + accept_len0: torch.Tensor | None, # [req_num0] + target_model_input1: ModelInput, # batch_size = padded_verify_batch_size1 + target_model_output1: ModelOutput, # logits: [padded_verify_batch_size1, vocab_size] + target_next_token_ids1: torch.Tensor, # [padded_verify_batch_size1] real_verify_rows1: int, - accept_len1: torch.Tensor | None, + accept_len1: torch.Tensor | None, # [req_num1] draft_step: int, ) -> SpecProposal: """Generate one proposal from two DP-overlapped decode microbatches.""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index 28ad442178..602aa086e0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -48,18 +48,18 @@ def fill_draft_model_kv_state_overlap( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, @@ -68,28 +68,28 @@ def propose_next( def propose_next_overlap( self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, real_verify_rows0: int, accept_len0: torch.Tensor | None, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + target_next_token_ids1: torch.Tensor, real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, ) -> EagleSpecProposal: return propose_next_dp_eagle_autoregressive_overlap( self, - main_model_input0, - main_model_output0, - next_token_ids0, + target_model_input0, + target_model_output0, + target_next_token_ids0, real_verify_rows0, accept_len0, - main_model_input1, - main_model_output1, - next_token_ids1, + target_model_input1, + target_model_output1, + target_next_token_ids1, real_verify_rows1, accept_len1, draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 00039d69e1..784e48aca8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -45,18 +45,18 @@ def fill_draft_model_kv_state_overlap( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, @@ -65,28 +65,28 @@ def propose_next( def propose_next_overlap( self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, real_verify_rows0: int, accept_len0: torch.Tensor | None, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + target_next_token_ids1: torch.Tensor, real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, ) -> EagleSpecProposal: return propose_next_dp_eagle_fixed_layout_overlap( self, - main_model_input0, - main_model_output0, - next_token_ids0, + target_model_input0, + target_model_output0, + target_next_token_ids0, real_verify_rows0, accept_len0, - main_model_input1, - main_model_output1, - next_token_ids1, + target_model_input1, + target_model_output1, + target_next_token_ids1, real_verify_rows1, accept_len1, draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 07ce70666b..2067855661 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -60,14 +60,14 @@ def pad_dp_step_mem_indexes( def propose_next_dp_eagle_autoregressive_overlap( proposer: BaseDpOverlapProposer, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, real_verify_rows0: int, accept_len0: torch.Tensor | None, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + target_next_token_ids1: torch.Tensor, real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, @@ -76,9 +76,9 @@ def propose_next_dp_eagle_autoregressive_overlap( """运行 DP EAGLE extend 后接单 token overlap decode 的 proposal 流程。""" verify_width = proposer.backend.max_draft_step + 1 - model_inputs = (main_model_input0, main_model_input1) - model_outputs = (main_model_output0, main_model_output1) - next_token_ids_by_batch = (next_token_ids0, next_token_ids1) + model_inputs = (target_model_input0, target_model_input1) + model_outputs = (target_model_output0, target_model_output1) + target_next_token_ids_by_batch = (target_next_token_ids0, target_next_token_ids1) real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) accept_lens_by_batch = (accept_len0, accept_len1) position_deltas_by_batch = tuple(model_input.b_position_delta for model_input in model_inputs) @@ -90,7 +90,7 @@ def propose_next_dp_eagle_autoregressive_overlap( for model_input, model_output, token_ids, real_verify_row_count, accept_len in zip( model_inputs, model_outputs, - next_token_ids_by_batch, + target_next_token_ids_by_batch, real_verify_row_counts, accept_lens_by_batch, ): @@ -114,7 +114,7 @@ def propose_next_dp_eagle_autoregressive_overlap( ) total_real_request_count = sum(real_request_counts) - proposal_token_ids = next_token_ids0.new_empty((total_real_request_count, draft_step)) + proposal_token_ids = target_next_token_ids0.new_empty((total_real_request_count, draft_step)) draft_model = proposer.backend.draft_models[0] extend_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) @@ -163,7 +163,7 @@ def propose_next_dp_eagle_autoregressive_overlap( model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids0.device, non_blocking=True) hold_mem_index = proposer.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX for step in range(1, draft_step): @@ -205,14 +205,14 @@ def propose_next_dp_eagle_autoregressive_overlap( def propose_next_dp_eagle_fixed_layout_overlap( proposer: BaseDpOverlapProposer, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, real_verify_rows0: int, accept_len0: torch.Tensor, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + target_next_token_ids1: torch.Tensor, real_verify_rows1: int, accept_len1: torch.Tensor, draft_step: int, @@ -221,31 +221,35 @@ def propose_next_dp_eagle_fixed_layout_overlap( """以 expanded verify-row layout 运行 decode,返回按真实请求压缩的 proposal。""" verify_width = proposer.backend.max_draft_step + 1 - model_inputs = (main_model_input0, main_model_input1) + model_inputs = (target_model_input0, target_model_input1) real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) real_request_counts = tuple(row_count // verify_width for row_count in real_verify_row_counts) request_capacities_by_batch = tuple(model_input.batch_size // verify_width for model_input in model_inputs) total_real_request_count = sum(real_request_counts) - proposal_token_ids = next_token_ids0.new_empty((total_real_request_count, draft_step)) + proposal_token_ids = target_next_token_ids0.new_empty((total_real_request_count, draft_step)) proposal_row_offsets = (0, real_request_counts[0]) accepted_tail_rows_by_batch = ( - torch.arange(0, main_model_input0.batch_size, verify_width, device=next_token_ids0.device) + accept_len0 - 1, - torch.arange(0, main_model_input1.batch_size, verify_width, device=next_token_ids1.device) + accept_len1 - 1, + torch.arange(0, target_model_input0.batch_size, verify_width, device=target_next_token_ids0.device) + + accept_len0 + - 1, + torch.arange(0, target_model_input1.batch_size, verify_width, device=target_next_token_ids1.device) + + accept_len1 + - 1, ) extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * draft_step) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids0.device, non_blocking=True) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids0.device, non_blocking=True) split = real_request_counts[0] * draft_step extra_mem_indexes_by_batch = ( extra_mem_indexes[:split], extra_mem_indexes[split:], ) - draft_token_ids_by_batch = [next_token_ids0, next_token_ids1] + draft_token_ids_by_batch = [target_next_token_ids0, target_next_token_ids1] draft_hiddens_by_batch = [ - main_model_output0.mtp_collector.spec_hidden, - main_model_output1.mtp_collector.spec_hidden, + target_model_output0.mtp_collector.spec_hidden, + target_model_output1.mtp_collector.spec_hidden, ] draft_model = proposer.backend.draft_models[0] hold_mem_index = proposer.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 61745960db..15ef075031 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -45,18 +45,18 @@ def fill_draft_model_kv_state_overlap( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, @@ -65,28 +65,28 @@ def propose_next( def propose_next_overlap( self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, real_verify_rows0: int, accept_len0: torch.Tensor | None, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + target_next_token_ids1: torch.Tensor, real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, ) -> EagleSpecProposal: return propose_next_dp_eagle_fixed_layout_overlap( self, - main_model_input0, - main_model_output0, - next_token_ids0, + target_model_input0, + target_model_output0, + target_next_token_ids0, real_verify_rows0, accept_len0, - main_model_input1, - main_model_output1, - next_token_ids1, + target_model_input1, + target_model_output1, + target_next_token_ids1, real_verify_rows1, accept_len1, draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 4ecb5590c3..bb0384b477 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -45,18 +45,18 @@ def fill_draft_model_kv_state_overlap( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: return propose_next_chained_mtp( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, @@ -64,28 +64,28 @@ def propose_next( def propose_next_overlap( self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, real_verify_rows0: int, accept_len0: torch.Tensor | None, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + target_next_token_ids1: torch.Tensor, real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, ) -> VanillaSpecProposal: return propose_next_dp_chained_mtp_overlap( self, - main_model_input0, - main_model_output0, - next_token_ids0, + target_model_input0, + target_model_output0, + target_next_token_ids0, real_verify_rows0, accept_len0, - main_model_input1, - main_model_output1, - next_token_ids1, + target_model_input1, + target_model_output1, + target_next_token_ids1, real_verify_rows1, accept_len1, draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py index cdd368bb6b..fa26f7c8c6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py @@ -44,34 +44,34 @@ def fill_dp_chained_mtp_draft_model_kv_state_overlap( def propose_next_dp_chained_mtp_overlap( proposer: BaseDpOverlapProposer, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, real_verify_rows0: int, accept_len0: torch.Tensor, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + target_next_token_ids1: torch.Tensor, real_verify_rows1: int, accept_len1: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: """为两个 DP microbatch 运行 decode,返回按真实请求压缩的 proposal。""" - model_inputs = (main_model_input0, main_model_input1) + model_inputs = (target_model_input0, target_model_input1) verify_width = proposer.backend.max_draft_step + 1 real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) real_request_counts = tuple(row_count // verify_width for row_count in real_verify_rows) accepted_tail_rows = ( - torch.arange(0, real_verify_rows0, verify_width, device=next_token_ids0.device) + accept_len0 - 1, - torch.arange(0, real_verify_rows1, verify_width, device=next_token_ids1.device) + accept_len1 - 1, + torch.arange(0, real_verify_rows0, verify_width, device=target_next_token_ids0.device) + accept_len0 - 1, + torch.arange(0, real_verify_rows1, verify_width, device=target_next_token_ids1.device) + accept_len1 - 1, ) - draft_token_ids = [next_token_ids0, next_token_ids1] + draft_token_ids = [target_next_token_ids0, target_next_token_ids1] draft_hiddens = [ - main_model_output0.mtp_collector.spec_hidden, - main_model_output1.mtp_collector.spec_hidden, + target_model_output0.mtp_collector.spec_hidden, + target_model_output1.mtp_collector.spec_hidden, ] - proposal_token_ids = next_token_ids0.new_empty((sum(real_request_counts), draft_step)) + proposal_token_ids = target_next_token_ids0.new_empty((sum(real_request_counts), draft_step)) request_offset = real_request_counts[0] for step in range(draft_step): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 301e2088d9..3d3d301c8d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -45,18 +45,18 @@ def fill_draft_model_kv_state_overlap( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: return propose_next_chained_mtp( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, @@ -64,28 +64,28 @@ def propose_next( def propose_next_overlap( self, - main_model_input0: ModelInput, - main_model_output0: ModelOutput, - next_token_ids0: torch.Tensor, + target_model_input0: ModelInput, + target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, real_verify_rows0: int, accept_len0: torch.Tensor | None, - main_model_input1: ModelInput, - main_model_output1: ModelOutput, - next_token_ids1: torch.Tensor, + target_model_input1: ModelInput, + target_model_output1: ModelOutput, + target_next_token_ids1: torch.Tensor, real_verify_rows1: int, accept_len1: torch.Tensor | None, draft_step: int, ) -> VanillaSpecProposal: return propose_next_dp_chained_mtp_overlap( self, - main_model_input0, - main_model_output0, - next_token_ids0, + target_model_input0, + target_model_output0, + target_next_token_ids0, real_verify_rows0, accept_len0, - main_model_input1, - main_model_output1, - next_token_ids1, + target_model_input1, + target_model_output1, + target_next_token_ids1, real_verify_rows1, accept_len1, draft_step, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py index 7187051b2a..951357110e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py @@ -25,18 +25,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py index 02c22af68a..052431f287 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py @@ -22,18 +22,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py index 3f94fb961a..53be483e49 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py @@ -22,18 +22,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py index 930977003e..ae51ab3219 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -22,18 +22,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: return propose_next_chained_mtp( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py index 559aaaa956..cb829c44fd 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -22,18 +22,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: return propose_next_chained_mtp( self, - main_model_input, - main_model_output, - next_token_ids, + target_model_input, + target_model_output, + target_next_token_ids, b_req_mtp_start_loc, draft_step, accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index c340d6c5cf..a0088f9570 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -111,17 +111,17 @@ def prepare_decode_model_input( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, + target_model_input: ModelInput, # batch_size = verify_batch_size + target_model_output: ModelOutput, # logits: [verify_batch_size, vocab_size] + target_next_token_ids: torch.Tensor, # [verify_batch_size] + b_req_mtp_start_loc: torch.Tensor, # [req_num] draft_step: int, - accept_len: Optional[torch.Tensor] = None, + accept_len: Optional[torch.Tensor] = None, # [req_num] ) -> SpecProposal: return self.proposer.propose_next( - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, accept_len=accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 1a144f9351..4b3d0ed6c6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -84,16 +84,16 @@ def fill_draft_model_kv_state( @abstractmethod def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, + target_model_input: ModelInput, # batch_size = verify_batch_size + target_model_output: ModelOutput, # logits: [verify_batch_size, vocab_size] + target_next_token_ids: torch.Tensor, # [verify_batch_size] + b_req_mtp_start_loc: torch.Tensor, # [req_num] draft_step: int, - accept_len: Optional[torch.Tensor] = None, + accept_len: Optional[torch.Tensor] = None, # [req_num] ) -> SpecProposal: """Generate candidate tokens after one target decode forward. - `main_model_input` contains the target verify rows, possibly compacted + `target_model_input` contains the target verify rows, possibly compacted by dynamic scheduling. `b_req_mtp_start_loc` identifies each logical request's first row. diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index e27838dde4..dff3f67309 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -47,9 +47,9 @@ def fill_draft_model_kv_state( @torch.no_grad() def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, @@ -62,15 +62,15 @@ def propose_next( accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() draft_input, extra_mem_indexes_cpu = build_parallel_block_draft_input( proposer=self, - main_model_input=main_model_input, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_next_token_ids=target_next_token_ids, accepted_tail_rows=accepted_tail_rows, request_count=request_count, ) extend_parallel_block_draft_kv_cache( proposer=self, - main_model_input=main_model_input, - target_hidden=main_model_output.mtp_collector.spec_hidden, + target_model_input=target_model_input, + target_hidden=target_model_output.mtp_collector.spec_hidden, ) draft_output = draft_model.forward(draft_input) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index ff3fb0da39..f7d22f18b8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -50,9 +50,9 @@ def fill_draft_model_kv_state( @torch.no_grad() def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, @@ -65,15 +65,15 @@ def propose_next( accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() draft_input, extra_mem_indexes_cpu = build_parallel_block_draft_input( proposer=self, - main_model_input=main_model_input, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_next_token_ids=target_next_token_ids, accepted_tail_rows=accepted_tail_rows, request_count=request_count, ) extend_parallel_block_draft_kv_cache( proposer=self, - main_model_input=main_model_input, - target_hidden=main_model_output.mtp_collector.spec_hidden, + target_model_input=target_model_input, + target_hidden=target_model_output.mtp_collector.spec_hidden, ) draft_output = draft_model.forward(draft_input) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py index a95775859e..544a762f7e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py @@ -33,18 +33,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( proposer=self, - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, accept_len=accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py index c1ca306b87..7c7d12f0f0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -22,18 +22,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( proposer=self, - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, accept_len=accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 976905e41c..5c7b9e0005 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -80,9 +80,9 @@ def prepare_eagle_verify_extend_input( def propose_next_eagle( proposer: BaseSpecProposer, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None, @@ -91,13 +91,13 @@ def propose_next_eagle( """运行 EAGLE extend 后接单 token decode 的通用 proposal 流程。""" request_count = int(b_req_mtp_start_loc.shape[0]) - proposal_token_ids = next_token_ids.new_empty((request_count, draft_step)) + proposal_token_ids = target_next_token_ids.new_empty((request_count, draft_step)) collect_schedule_scores = proposer.enable_dynmaic_mtp schedule_scores = ( torch.zeros( (request_count, draft_step), dtype=torch.float32, - device=next_token_ids.device, + device=target_next_token_ids.device, ) if collect_schedule_scores else None @@ -111,13 +111,13 @@ def propose_next_eagle( accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() draft_model = proposer.backend.draft_models[0] - position_delta = main_model_input.b_position_delta + position_delta = target_model_input.b_position_delta prepare_eagle_verify_extend_input( - model_input=main_model_input, - input_ids=next_token_ids, - target_hidden=main_model_output.mtp_collector.spec_hidden, + model_input=target_model_input, + input_ids=target_next_token_ids, + target_hidden=target_model_output.mtp_collector.spec_hidden, ) - extend_output = draft_model.forward(main_model_input) + extend_output = draft_model.forward(target_model_input) accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) if collect_schedule_scores: @@ -144,13 +144,13 @@ def propose_next_eagle( ) extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(request_count * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) - draft_seq_lens = main_model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1 - max_kv_seq_len = main_model_input.max_kv_seq_len - draft_input = copy.copy(main_model_input) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids.device, non_blocking=True) + draft_seq_lens = target_model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1 + max_kv_seq_len = target_model_input.max_kv_seq_len + draft_input = copy.copy(target_model_input) draft_input.is_prefill = False draft_input.batch_size = request_count - draft_input.b_req_idx = main_model_input.b_req_idx.index_select(0, accepted_tail_rows) + draft_input.b_req_idx = target_model_input.b_req_idx.index_select(0, accepted_tail_rows) draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) draft_input.b_seq_len = draft_seq_lens draft_input.b_position_delta = ( diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index 7cd6ffc701..d5bab30594 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -22,18 +22,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: return propose_next_eagle( proposer=self, - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, accept_len=accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py index 2a4a5e784e..20d103d07a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py @@ -30,28 +30,28 @@ def fill_parallel_block_draft_model_kv_state( def extend_parallel_block_draft_kv_cache( proposer: BaseSpecProposer, - main_model_input: ModelInput, + target_model_input: ModelInput, target_hidden: torch.Tensor, ) -> None: """提交本轮 target verify hidden,扩展 parallel-block drafter KV。""" - main_model_input.total_token_num = main_model_input.batch_size - main_model_input.prefix_total_token_num = 0 - main_model_input.is_prefill = True - main_model_input.b_ready_cache_len = main_model_input.b_seq_len - 1 - main_model_input.b_prefill_start_loc = torch.arange( - main_model_input.batch_size, + target_model_input.total_token_num = target_model_input.batch_size + target_model_input.prefix_total_token_num = 0 + target_model_input.is_prefill = True + target_model_input.b_ready_cache_len = target_model_input.b_seq_len - 1 + target_model_input.b_prefill_start_loc = torch.arange( + target_model_input.batch_size, dtype=torch.int32, device=target_hidden.device, ) - main_model_input.mtp_draft_input_hiddens = target_hidden - proposer.backend.draft_models[0].forward(main_model_input) + target_model_input.mtp_draft_input_hiddens = target_hidden + proposer.backend.draft_models[0].forward(target_model_input) def build_parallel_block_draft_input( proposer: BaseSpecProposer, - main_model_input: ModelInput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_next_token_ids: torch.Tensor, accepted_tail_rows: torch.Tensor, request_count: int, ): @@ -61,36 +61,38 @@ def build_parallel_block_draft_input( block_size = int(draft_model.block_size) extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(request_count * block_size) - block_input_ids = next_token_ids.new_full( + block_input_ids = target_next_token_ids.new_full( (request_count * block_size,), fill_value=draft_model.mask_token_id, ) - block_input_ids[::block_size] = next_token_ids.index_select(0, accepted_tail_rows) + block_input_ids[::block_size] = target_next_token_ids.index_select(0, accepted_tail_rows) block_offsets = torch.arange( block_size, - dtype=main_model_input.b_seq_len.dtype, - device=next_token_ids.device, + dtype=target_model_input.b_seq_len.dtype, + device=target_next_token_ids.device, ) - draft_input = copy.copy(main_model_input) + draft_input = copy.copy(target_model_input) draft_input.input_ids = block_input_ids draft_input.total_token_num = draft_input.input_ids.shape[0] draft_input.batch_size = draft_input.total_token_num draft_input.max_q_seq_len = 1 - draft_input.max_kv_seq_len = main_model_input.max_kv_seq_len + block_size + draft_input.max_kv_seq_len = target_model_input.max_kv_seq_len + block_size draft_input.b_req_idx = ( - main_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() + target_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() ) draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) draft_input.b_seq_len = ( - (main_model_input.b_seq_len.index_select(0, accepted_tail_rows)[:, None] + block_offsets[None, :] + 1) + (target_model_input.b_seq_len.index_select(0, accepted_tail_rows)[:, None] + block_offsets[None, :] + 1) .reshape(-1) .contiguous() ) draft_input.b_position_delta = ( - main_model_input.b_position_delta.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() + target_model_input.b_position_delta.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() ) - draft_input.mem_indexes = extra_mem_indexes_cpu.to(device=next_token_ids.device, non_blocking=True) + draft_input.mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids.device, non_blocking=True) draft_input.b_mark_shared_group = torch.zeros_like(draft_input.b_req_idx) draft_input.b_mark_shared_group[block_size - 1 :: block_size] = block_size empty_multimodal_params = {"images": [], "audios": []} diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index 58021a295e..da85a0defe 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -27,18 +27,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: return propose_next_chained_mtp( proposer=self, - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, accept_len=accept_len, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py index 834f978b39..ae6f00b607 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py @@ -42,9 +42,9 @@ def fill_chained_mtp_draft_model_kv_state( def propose_next_chained_mtp( proposer: BaseSpecProposer, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor, @@ -53,14 +53,14 @@ def propose_next_chained_mtp( request_count = int(b_req_mtp_start_loc.shape[0]) accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_token_ids = next_token_ids - draft_hidden = main_model_output.mtp_collector.spec_hidden - proposal_token_ids = next_token_ids.new_empty((request_count, draft_step)) + draft_token_ids = target_next_token_ids + draft_hidden = target_model_output.mtp_collector.spec_hidden + proposal_token_ids = target_next_token_ids.new_empty((request_count, draft_step)) schedule_scores = ( torch.empty( (request_count, draft_step), dtype=torch.float32, - device=next_token_ids.device, + device=target_next_token_ids.device, ) if proposer.enable_dynmaic_mtp else None @@ -68,9 +68,9 @@ def propose_next_chained_mtp( for step in range(draft_step): draft_model = proposer.backend.draft_models[step] - main_model_input.input_ids = draft_token_ids - main_model_input.mtp_draft_input_hiddens = draft_hidden - draft_output = draft_model.forward(main_model_input) + target_model_input.input_ids = draft_token_ids + target_model_input.mtp_draft_input_hiddens = draft_hidden + draft_output = draft_model.forward(target_model_input) draft_hidden = draft_output.mtp_collector.spec_hidden if proposer.enable_dynmaic_mtp: draft_token_ids, draft_token_probs = proposer.backend._gen_argmax_token_ids_and_prob(draft_output) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index 559dd0ef8d..18af6acbef 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -27,18 +27,18 @@ def fill_draft_model_kv_state( def propose_next( self, - main_model_input: ModelInput, - main_model_output: ModelOutput, - next_token_ids: torch.Tensor, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: return propose_next_chained_mtp( proposer=self, - main_model_input=main_model_input, - main_model_output=main_model_output, - next_token_ids=next_token_ids, + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=draft_step, accept_len=accept_len, diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index dbabd15bdc..ac7efc2b45 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -22,7 +22,7 @@ def __init__(self): def propose_next(self, **kwargs): self.propose_args = kwargs - token_ids = kwargs["next_token_ids"].new_zeros((kwargs["b_req_mtp_start_loc"].shape[0], 7)) + token_ids = kwargs["target_next_token_ids"].new_zeros((kwargs["b_req_mtp_start_loc"].shape[0], 7)) return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([123], dtype=torch.int32))], @@ -31,7 +31,7 @@ def propose_next(self, **kwargs): def propose_next_overlap(self, **kwargs): self.propose_overlap_args = kwargs request_count = (kwargs["real_verify_rows0"] + kwargs["real_verify_rows1"]) // 8 - token_ids = kwargs["next_token_ids0"].new_zeros((request_count, 7)) + token_ids = kwargs["target_next_token_ids0"].new_zeros((request_count, 7)) return SpecProposal( token_ids=token_ids, extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([456], dtype=torch.int32))], @@ -129,8 +129,8 @@ def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(monkeypatch): ) propose_args = backend.spec_engine.propose_args - assert propose_args["next_token_ids"].shape == (16,) - assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) + assert propose_args["target_next_token_ids"].shape == (16,) + assert torch.equal(propose_args["target_next_token_ids"][:8], next_token_ids) assert torch.equal(propose_args["b_req_mtp_start_loc"], torch.tensor([0, 8], dtype=torch.int32)) assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) assert scatter_args["proposal"].token_ids.shape == (1, 7) @@ -160,8 +160,8 @@ def test_dp_vanilla_uses_dp_engine_proposer(monkeypatch): ) propose_args = backend.spec_engine.propose_args - assert propose_args["next_token_ids"].shape == (16,) - assert torch.equal(propose_args["next_token_ids"][:8], next_token_ids) + assert propose_args["target_next_token_ids"].shape == (16,) + assert torch.equal(propose_args["target_next_token_ids"][:8], next_token_ids) assert scatter_args["proposal"].token_ids.shape == (8, 7) assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([123], dtype=torch.int32)) @@ -203,8 +203,8 @@ def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(monkeypat ) propose_args = backend.dp_overlap_spec_engine.propose_overlap_args - assert propose_args["next_token_ids0"].shape == (16,) - assert propose_args["next_token_ids1"].shape == (16,) + assert propose_args["target_next_token_ids0"].shape == (16,) + assert propose_args["target_next_token_ids1"].shape == (16,) assert propose_args["real_verify_rows0"] == 8 assert propose_args["real_verify_rows1"] == 16 assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) @@ -245,8 +245,8 @@ def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(monkeypatch) ) propose_args = backend.dp_overlap_spec_engine.propose_overlap_args - assert propose_args["next_token_ids0"].shape == (16,) - assert propose_args["next_token_ids1"].shape == (16,) + assert propose_args["target_next_token_ids0"].shape == (16,) + assert propose_args["target_next_token_ids1"].shape == (16,) assert propose_args["real_verify_rows0"] == 8 assert propose_args["real_verify_rows1"] == 16 assert torch.equal(propose_args["accept_len0"], torch.tensor([2], dtype=torch.int32)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 22a9752e79..40f11d218a 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -85,18 +85,18 @@ def test_overlap_eagle_keeps_fixed_verify_layout(monkeypatch): model_input1 = _target_input(batch_size=6) proposal = proposer.propose_next_overlap( - main_model_input0=model_input0, - main_model_output0=ModelOutput( + target_model_input0=model_input0, + target_model_output0=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), - next_token_ids0=torch.arange(6, dtype=torch.int64), + target_next_token_ids0=torch.arange(6, dtype=torch.int64), real_verify_rows0=3, accept_len0=torch.tensor([2, 1], dtype=torch.int32), - main_model_input1=model_input1, - main_model_output1=ModelOutput( + target_model_input1=model_input1, + target_model_output1=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), - next_token_ids1=torch.arange(10, 16, dtype=torch.int64), + target_next_token_ids1=torch.arange(10, 16, dtype=torch.int64), real_verify_rows1=6, accept_len1=torch.tensor([1, 3], dtype=torch.int32), draft_step=2, @@ -137,18 +137,18 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): model_input1 = _target_input(batch_size=6) proposal = proposer.propose_next_overlap( - main_model_input0=model_input0, - main_model_output0=ModelOutput( + target_model_input0=model_input0, + target_model_output0=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), - next_token_ids0=torch.arange(6, dtype=torch.int64), + target_next_token_ids0=torch.arange(6, dtype=torch.int64), real_verify_rows0=3, accept_len0=torch.tensor([2, 1], dtype=torch.int32), - main_model_input1=model_input1, - main_model_output1=ModelOutput( + target_model_input1=model_input1, + target_model_output1=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), - next_token_ids1=torch.arange(10, 16, dtype=torch.int64), + target_next_token_ids1=torch.arange(10, 16, dtype=torch.int64), real_verify_rows1=6, accept_len1=torch.tensor([1, 3], dtype=torch.int32), draft_step=2, diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 7458cd6d2c..c4a10cf5df 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -359,9 +359,9 @@ def test_eagle_proposer_skips_draft_forward_for_zero_steps(): next_token_ids = torch.tensor([10, 11], dtype=torch.int64) proposal = proposer.propose_next( - main_model_input=None, - main_model_output=None, - next_token_ids=next_token_ids, + target_model_input=None, + target_model_output=None, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=torch.tensor([0, 1], dtype=torch.int32), draft_step=0, ) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py index 0cb075b80f..60885ed0e2 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -35,20 +35,20 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): model_input1 = SimpleNamespace(batch_size=6) proposal = proposer.propose_next_overlap( - main_model_input0=model_input0, - main_model_output0=ModelOutput( + target_model_input0=model_input0, + target_model_output0=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), ), - next_token_ids0=torch.tensor([10, 11, 0, 0, 0, 0], dtype=torch.int64), + target_next_token_ids0=torch.tensor([10, 11, 0, 0, 0, 0], dtype=torch.int64), real_verify_rows0=3, accept_len0=torch.tensor([2], dtype=torch.int32), - main_model_input1=model_input1, - main_model_output1=ModelOutput( + target_model_input1=model_input1, + target_model_output1=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), ), - next_token_ids1=torch.tensor([20, 21, 22, 0, 0, 0], dtype=torch.int64), + target_next_token_ids1=torch.tensor([20, 21, 22, 0, 0, 0], dtype=torch.int64), real_verify_rows1=3, accept_len1=torch.tensor([1], dtype=torch.int32), draft_step=2, diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index b480c13db8..e17f946f4c 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -249,11 +249,11 @@ def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(monkeypatch) monkeypatch.setattr(dflash_module, "extend_parallel_block_draft_kv_cache", lambda **_: None) proposal = proposer.propose_next( - main_model_input=SimpleNamespace(), - main_model_output=SimpleNamespace( + target_model_input=SimpleNamespace(), + target_model_output=SimpleNamespace( mtp_collector=SimpleNamespace(spec_hidden=torch.empty(verify_row_count, 1)), ), - next_token_ids=torch.arange(verify_row_count), + target_next_token_ids=torch.arange(verify_row_count), b_req_mtp_start_loc=torch.tensor([0, 3]), draft_step=2, accept_len=torch.tensor([1, 1]), @@ -328,8 +328,8 @@ def test_dflash_expands_position_delta_with_request_block_rows(monkeypatch): draft_input, _ = build_parallel_block_draft_input( proposer=proposer, - main_model_input=model_input, - next_token_ids=torch.arange(5, dtype=torch.int64), + target_model_input=model_input, + target_next_token_ids=torch.arange(5, dtype=torch.int64), accepted_tail_rows=torch.tensor([1, 4]), request_count=2, ) From 57436247991787c05278f0450a341fdbdbe0ea80 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 00:38:26 +0000 Subject: [PATCH 062/103] refactor: skip vanilla no-att draft state fill --- .../mtp_speculative/proposers/vanilla_no_att.py | 8 +------- 1 file changed, 1 insertion(+), 7 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index da85a0defe..cbe3d2a6ac 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -4,7 +4,6 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -18,12 +17,7 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - ) + pass def propose_next( self, From 46c8bdaa58099465fdbcf728f9f741fccc81de80 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 03:11:26 +0000 Subject: [PATCH 063/103] refactor: build MTP attention groups in backends --- lightllm/common/basemodel/attention/fa3/fp.py | 7 +- .../common/basemodel/attention/fa3/mla.py | 7 +- .../common/basemodel/attention/triton/fp.py | 12 ++- lightllm/common/basemodel/basemodel.py | 15 +--- lightllm/common/basemodel/batch_objs.py | 9 +- .../triton_kernel/dynamic_mtp_utils.py | 63 +------------ .../basemodel/triton_kernel/fa3_utils.py | 57 +++++++----- .../basemodel/triton_kernel/mtp_utils.py | 88 +++++++++++++++++++ .../generic_padded_pre_process.py | 5 -- .../mode_backend/generic_pre_process.py | 42 --------- .../dp_overlap_proposers/eagle_utils.py | 1 - .../mtp_speculative/proposers/eagle_utils.py | 1 - .../proposers/parallel_block_utils.py | 2 - .../common/basemodel/test_model_input.py | 48 ++-------- .../triton_kernel/test_dynamic_mtp_utils.py | 13 +-- .../basemodel/triton_kernel/test_fa3_utils.py | 18 ++-- .../basemodel/triton_kernel/test_mtp_utils.py | 18 ++++ .../mode_backend/test_generic_pre_process.py | 26 +----- unit_tests/utils/test_speculative_utils.py | 41 +++++++++ 19 files changed, 234 insertions(+), 239 deletions(-) diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index be190a05bc..ea11b8307d 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -9,6 +9,7 @@ page_table_copy, ) from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor +from lightllm.common.basemodel.triton_kernel.mtp_utils import build_mtp_shared_group_markers from lightllm.utils.envs_utils import get_env_start_args @@ -148,10 +149,14 @@ def init_state(self): self._init_page_table(b_att_req_idx) def _init_dynamic_spec_verify_state(self, draft_step: int) -> torch.Tensor: + b_mark_mtp_shared_group = build_mtp_shared_group_markers( + self.infer_state.b_req_idx, + hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, + ) b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len = build_dynamic_spec_fa3_decode_params( b_req_idx=self.infer_state.b_req_idx, b_seq_len=self.infer_state.b_seq_len, - b_mark_shared_group=self.infer_state.b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, att_batch_size=self.infer_state.batch_size, hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, ) diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index 8805c6cae4..3f4a9170a6 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -6,6 +6,7 @@ from lightllm.utils.sgl_utils import flash_attn_with_kvcache from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_spec_fa3_decode_params, page_table_copy from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor +from lightllm.common.basemodel.triton_kernel.mtp_utils import build_mtp_shared_group_markers from lightllm.utils.sgl_utils import flash_attn_varlen_func from lightllm.utils.envs_utils import get_env_start_args @@ -128,10 +129,14 @@ def init_state(self): self._init_page_table(b_att_req_idx) def _init_dynamic_spec_verify_state(self, draft_step: int) -> torch.Tensor: + b_mark_mtp_shared_group = build_mtp_shared_group_markers( + self.infer_state.b_req_idx, + hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, + ) b_q_seq_len, b_kv_seq_len, b_att_req_idx, self.b_att_seq_len = build_dynamic_spec_fa3_decode_params( b_req_idx=self.infer_state.b_req_idx, b_seq_len=self.infer_state.b_seq_len, - b_mark_shared_group=self.infer_state.b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, att_batch_size=self.infer_state.batch_size, hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, ) diff --git a/lightllm/common/basemodel/attention/triton/fp.py b/lightllm/common/basemodel/attention/triton/fp.py index 9a57f2dc28..e7ce66c774 100644 --- a/lightllm/common/basemodel/attention/triton/fp.py +++ b/lightllm/common/basemodel/attention/triton/fp.py @@ -2,6 +2,7 @@ import torch from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl from typing import Optional +from lightllm.common.basemodel.triton_kernel.mtp_utils import build_mtp_shared_group_markers class TritonAttBackend(BaseAttBackend): @@ -93,8 +94,15 @@ def _nomarl_prefill_att( @dataclasses.dataclass class TritonDecodeAttState(BaseDecodeAttState): + b_mark_mtp_shared_group: torch.Tensor = None + def init_state(self): - pass + draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) + if draft_step > 0: + self.b_mark_mtp_shared_group = build_mtp_shared_group_markers( + self.infer_state.b_req_idx, + hold_req_id=self.backend.model.req_manager.HOLD_REQUEST_ID, + ) def copy_for_decode_cuda_graph(self, new_state: "TritonDecodeAttState"): super().copy_for_decode_cuda_graph(new_state) @@ -229,7 +237,7 @@ def _spec_decode_gqa_att( Req_to_tokens=self.infer_state.req_manager.req_to_token_indexs, B_req_idx=self.infer_state.b_req_idx, b_seq_len=self.infer_state.b_seq_len, - b_mark_shared_group=self.infer_state.b_mark_shared_group, + b_mark_shared_group=self.b_mark_mtp_shared_group, alloc_tensor_func=alloc_func, ) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index bce508adb8..0df46554e1 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -37,7 +37,6 @@ set_model_init_status, enable_diverse_mode_gqa_decode_fast_kernel, enable_full_att_decode_tune, - enable_triton_mtp_kernel, ) from lightllm.common.triton_utils.autotuner import Autotuner from lightllm.utils.infer_utils import post_empty_cache @@ -401,12 +400,9 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) infer_state.b_ready_cache_len = model_input.b_ready_cache_len else: infer_state.b_ready_cache_len = torch.zeros_like(input=infer_state.b_seq_len) - else: - if enable_diverse_mode_gqa_decode_fast_kernel(): - infer_state.b_shared_seq_len = model_input.b_shared_seq_len - infer_state.b_mark_shared_group = model_input.b_mark_shared_group - elif self.args.mtp_dynamic_verify or enable_triton_mtp_kernel(): - infer_state.b_mark_shared_group = model_input.b_mark_shared_group + elif enable_diverse_mode_gqa_decode_fast_kernel(): + infer_state.b_shared_seq_len = model_input.b_shared_seq_len + infer_state.b_mark_shared_group = model_input.b_mark_shared_group infer_state.multimodal_params = model_input.multimodal_params @@ -477,11 +473,6 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.b_mark_shared_group = F.pad( new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=1 ) - elif self.args.mtp_dynamic_verify or enable_triton_mtp_kernel(): - assert new_model_input.b_mark_shared_group is not None - new_model_input.b_mark_shared_group = F.pad( - new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=0 - ) # 特殊模型,特殊模式的特殊变量的特殊 padding if new_model_input.mtp_draft_input_hiddens is not None: diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 169edd6661..68637874ef 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -4,11 +4,7 @@ import torch -from lightllm.utils.envs_utils import ( - enable_diverse_mode_gqa_decode_fast_kernel, - enable_triton_mtp_kernel, - get_env_start_args, -) +from lightllm.utils.envs_utils import enable_diverse_mode_gqa_decode_fast_kernel from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor @@ -116,9 +112,6 @@ def _ensure_decode_group_metadata(self): self.b_mark_shared_group = torch.ones_like(self.b_req_idx, dtype=torch.int32) if self.b_shared_seq_len is None: self.b_shared_seq_len = torch.zeros_like(self.b_req_idx, dtype=torch.int32) - elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): - if self.b_mark_shared_group is None: - self.b_mark_shared_group = torch.ones_like(self.b_req_idx, dtype=torch.int32) @dataclass diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py index 70a5eaf700..0ca62dbd11 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -8,7 +8,6 @@ from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.basemodel.mtp_manager import MtpManager -from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size # 动态 verify 行选择。 @@ -138,37 +137,6 @@ def _fwd_kernel_compact_dynamic_mtp_model_input( return -@triton.jit -def _fwd_kernel_rebuild_trimmed_mtp_b_mark_shared_group( - b_req_idx, - out_b_mark_shared_group, - batch_size, - max_batch_shared_group_size: tl.constexpr, - MAX_RUN_SCAN: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - offsets = tl.arange(0, BLOCK_SIZE) - mask = offsets < batch_size - cur_req_idx = tl.load(b_req_idx + offsets, mask=mask, other=-1) - - prev_same_count = tl.full((BLOCK_SIZE,), 0, tl.int32) - for scan_offset in tl.static_range(1, MAX_RUN_SCAN + 1): - prev_offsets = offsets - scan_offset - prev_mask = mask & (prev_offsets >= 0) - prev_req_idx = tl.load(b_req_idx + prev_offsets, mask=prev_mask, other=-2) - prev_same_count += tl.where(prev_mask & (prev_req_idx == cur_req_idx), 1, 0) - - next_offsets = offsets + 1 - next_req_idx = tl.load(b_req_idx + next_offsets, mask=next_offsets < batch_size, other=-2) - group_pos = prev_same_count % max_batch_shared_group_size - is_group_end = mask & ( - (next_offsets == batch_size) | (next_req_idx != cur_req_idx) | (group_pos == max_batch_shared_group_size - 1) - ) - mark_value = tl.where(is_group_end, group_pos + 1, 0) - tl.store(out_b_mark_shared_group + offsets, mark_value, mask=mask) - return - - @triton.jit def _fwd_kernel_pack_selected_rows_2d( src, @@ -232,35 +200,10 @@ def _pack_selected_hidden( return dst -def _rebuild_mtp_group_markers(b_req_idx: torch.Tensor, max_request_rows: int) -> torch.Tensor: - assert b_req_idx.is_cuda - batch_size = b_req_idx.shape[0] - max_batch_shared_group_size = int(get_diverse_max_batch_shared_group_size()) - assert max_batch_shared_group_size > 0 - assert max_request_rows > 0 - if batch_size == 0: - return torch.empty((0,), dtype=torch.int32, device=b_req_idx.device) - - b_mark_shared_group = torch.empty((batch_size,), dtype=torch.int32, device=b_req_idx.device) - BLOCK_SIZE = triton.next_power_of_2(batch_size) - _fwd_kernel_rebuild_trimmed_mtp_b_mark_shared_group[(1,)]( - b_req_idx=b_req_idx, - out_b_mark_shared_group=b_mark_shared_group, - batch_size=batch_size, - max_batch_shared_group_size=max_batch_shared_group_size, - MAX_RUN_SCAN=max_request_rows - 1, - BLOCK_SIZE=BLOCK_SIZE, - num_warps=8, - num_stages=1, - ) - return b_mark_shared_group - - def _compact_decode_model_input( model_input: ModelInput, selected_row_mask: torch.Tensor, dynamic_batch_size: int, - max_draft_step: int, ) -> ModelInput: assert not model_input.is_prefill assert selected_row_mask.is_cuda @@ -338,10 +281,7 @@ def _compact_decode_model_input( model_input.b_seq_len = out_b_seq_len model_input.b_position_delta = out_b_position_delta model_input.b_shared_seq_len = out_b_shared_seq_len - model_input.b_mark_shared_group = _rebuild_mtp_group_markers( - out_b_req_idx, - max_request_rows=max_draft_step + 1, - ) + model_input.b_mark_shared_group = None if model_input.mtp_draft_input_hiddens is not None: assert model_input.mtp_draft_input_hiddens.is_cuda @@ -391,7 +331,6 @@ def prepare_dynamic_mtp_model_input( model_input=model_input, selected_row_mask=selected_row_mask, dynamic_batch_size=dynamic_batch_size, - max_draft_step=max_draft_step, ) # Decode and draft-cache commit use the compacted b_position_delta, so # placeholder multimodal metadata only needs to keep ModelInput shapes diff --git a/lightllm/common/basemodel/triton_kernel/fa3_utils.py b/lightllm/common/basemodel/triton_kernel/fa3_utils.py index 520e8a7859..3d04558273 100644 --- a/lightllm/common/basemodel/triton_kernel/fa3_utils.py +++ b/lightllm/common/basemodel/triton_kernel/fa3_utils.py @@ -66,7 +66,7 @@ def page_table_copy( def _build_dynamic_spec_fa3_decode_params_kernel( b_req_idx, b_seq_len, - b_mark_shared_group, + b_mark_mtp_shared_group, out_b_q_seq_len, out_b_kv_seq_len, out_b_att_req_idx, @@ -78,7 +78,7 @@ def _build_dynamic_spec_fa3_decode_params_kernel( offsets = tl.arange(0, BLOCK_SIZE) mask = offsets < batch_size - mark = tl.load(b_mark_shared_group + offsets, mask=mask, other=0) + mark = tl.load(b_mark_mtp_shared_group + offsets, mask=mask, other=0) # A positive mark closes a query group and stores the number of rows in it. is_group_end = mask & (mark > 0) dst_pos = tl.cumsum(tl.where(is_group_end, 1, 0), axis=0) - 1 @@ -99,7 +99,7 @@ def _build_dynamic_spec_fa3_decode_params_kernel( @triton.jit def _count_dynamic_spec_fa3_decode_params_kernel( - b_mark_shared_group, + b_mark_mtp_shared_group, out_block_counts, batch_size, BLOCK_SIZE: tl.constexpr, @@ -108,7 +108,7 @@ def _count_dynamic_spec_fa3_decode_params_kernel( offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < batch_size - mark = tl.load(b_mark_shared_group + offsets, mask=mask, other=0) + mark = tl.load(b_mark_mtp_shared_group + offsets, mask=mask, other=0) block_count = tl.sum(tl.where(mask & (mark > 0), 1, 0), axis=0) tl.store(out_block_counts + block_id, block_count) @@ -117,7 +117,7 @@ def _count_dynamic_spec_fa3_decode_params_kernel( def _compact_dynamic_spec_fa3_decode_params_kernel( b_req_idx, b_seq_len, - b_mark_shared_group, + b_mark_mtp_shared_group, block_counts, block_offsets, out_b_q_seq_len, @@ -132,7 +132,7 @@ def _compact_dynamic_spec_fa3_decode_params_kernel( offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < batch_size - mark = tl.load(b_mark_shared_group + offsets, mask=mask, other=0) + mark = tl.load(b_mark_mtp_shared_group + offsets, mask=mask, other=0) # A positive mark closes a query group and stores the number of rows in it. is_group_end = mask & (mark > 0) local_pos = tl.cumsum(tl.where(is_group_end, 1, 0), axis=0) - 1 @@ -154,35 +154,50 @@ def _compact_dynamic_spec_fa3_decode_params_kernel( def build_dynamic_spec_fa3_decode_params( b_req_idx: torch.Tensor, b_seq_len: torch.Tensor, - b_mark_shared_group: torch.Tensor, + b_mark_mtp_shared_group: torch.Tensor, att_batch_size: int, hold_req_id: int, ): """Convert dynamic speculative-verify rows to one metadata row per FA3 sequence. Query rows belonging to the same attention sequence are consecutive. - ``b_mark_shared_group`` is zero inside a group; its final row stores the + ``b_mark_mtp_shared_group`` is zero inside a group; its final row stores the number of query rows in that group. For example, let ``H`` denote ``hold_req_id`` and suppose:: b_req_idx = [ 7, 7, 7, 4, 9, 9, H, H] - b_seq_len = [12, 13, 14, 8, 20, 21, 0, 0] - b_mark_shared_group = [ 0, 0, 3, 1, 0, 2, 0, 0] + b_seq_len = [12, 13, 14, 8, 20, 21, 2, 2] + b_mark_mtp_shared_group = [ 0, 0, 3, 1, 0, 2, 1, 1] - The positive marks close three attention sequences:: + The positive marks close five attention sequences. HOLD rows remain + independent one-row groups:: request 7 -> query length 3, final KV length 14 request 4 -> query length 1, final KV length 8 request 9 -> query length 2, final KV length 21 + HOLD row 6 -> query length 1 + HOLD row 7 -> query length 1 The operator retains each group's final row, compacts those rows to the front, and pads the unused tail to ``att_batch_size``:: - b_q_seq_len = [ 3, 1, 2, 0, 0, 0, 0, 0] - b_kv_seq_len = [14, 8, 21, 0, 0, 0, 0, 0] - b_att_req_idx = [7, 4, 9, H, H, H, H, H] - b_att_seq_len = [14, 8, 21, 0, 0, 0, 0, 0] + # executable metadata + # HOLD groups padding + # | | | | | + b_q_seq_len = [ 3, 1, 2, 1, 1, 0, 0, 0] + b_kv_seq_len = [14, 8, 21, 2, 2, 0, 0, 0] + b_att_req_idx = [ 7, 4, 9, H, H, H, H, H] + b_att_seq_len = [14, 8, 21, 2, 2, 0, 0, 0] + + The first two ``H`` entries are input HOLD rows. Each remains an + executable one-query dummy sequence with a safe KV length, so FA3 consumes + every physical query row without merging adjacent HOLD requests. The final + three ``H`` entries only pad the metadata tensors to ``att_batch_size``; + their zero query and KV lengths prevent them from describing attention + work. Both kinds use ``hold_req_id`` because every ``b_att_req_idx`` value + must remain a valid page-table row. Their sequence lengths, rather than the + request id, distinguish executable HOLD groups from metadata padding. Thus ``b_att_req_idx`` changes from one request id per query row to one id per FA3 sequence; it selects that sequence's page-table row. A request may @@ -190,8 +205,8 @@ def build_dynamic_spec_fa3_decode_params( multiple attention groups. Fixed output shapes let one CUDA Graph replay different speculative verify layouts. """ - assert b_req_idx.is_cuda and b_seq_len.is_cuda and b_mark_shared_group.is_cuda - assert b_req_idx.shape == b_seq_len.shape == b_mark_shared_group.shape + assert b_req_idx.is_cuda and b_seq_len.is_cuda and b_mark_mtp_shared_group.is_cuda + assert b_req_idx.shape == b_seq_len.shape == b_mark_mtp_shared_group.shape assert b_req_idx.shape[0] == att_batch_size assert att_batch_size > 0 @@ -204,7 +219,7 @@ def build_dynamic_spec_fa3_decode_params( _build_dynamic_spec_fa3_decode_params_kernel[(1,)]( b_req_idx=b_req_idx, b_seq_len=b_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, out_b_q_seq_len=b_q_seq_len, out_b_kv_seq_len=b_kv_seq_len, out_b_att_req_idx=b_att_req_idx, @@ -224,10 +239,10 @@ def build_dynamic_spec_fa3_decode_params( block_size = _DYNAMIC_SPEC_FA3_COMPACT_BLOCK_SIZE grid = (triton.cdiv(att_batch_size, block_size),) - block_counts = torch.empty((grid[0],), dtype=torch.int32, device=b_mark_shared_group.device) + block_counts = torch.empty((grid[0],), dtype=torch.int32, device=b_mark_mtp_shared_group.device) _count_dynamic_spec_fa3_decode_params_kernel[grid]( - b_mark_shared_group=b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, out_block_counts=block_counts, batch_size=att_batch_size, BLOCK_SIZE=block_size, @@ -239,7 +254,7 @@ def build_dynamic_spec_fa3_decode_params( _compact_dynamic_spec_fa3_decode_params_kernel[grid]( b_req_idx=b_req_idx, b_seq_len=b_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, block_counts=block_counts, block_offsets=block_offsets, out_b_q_seq_len=b_q_seq_len, diff --git a/lightllm/common/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index f4297866de..7e943f2925 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -6,6 +6,94 @@ import triton.language as tl import torch +from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size + + +@triton.jit +def _fwd_kernel_build_mtp_shared_group_markers( + b_req_idx, + b_mark_mtp_shared_group, + batch_size, + hold_req_id, + MAX_GROUP_SIZE: tl.constexpr, + SCAN_BLOCK_SIZE: tl.constexpr, +): + current_row = tl.program_id(axis=0) + current_req_idx = tl.load(b_req_idx + current_row) + + # HOLD rows never participate in request grouping. Each one directly forms + # an independent one-row group, even when multiple HOLD rows are adjacent. + has_next_row = current_row + 1 < batch_size + next_req_idx = tl.load( + b_req_idx + current_row + 1, + mask=has_next_row, + other=-1, + ) + is_hold_row = current_req_idx == hold_req_id + if is_hold_row: + tl.store(b_mark_mtp_shared_group + current_row, 1) + return + + # A normal request only needs work on the final row of its consecutive run. + is_request_run_end = (~has_next_row) | (next_req_idx != current_req_idx) + if not is_request_run_end: + return + + # ModelInput keeps all MTP rows of one real request together and emits each + # request only once. Count the rows matching the current request, then use + # that count to recover the request's first row. + backward_offsets = tl.arange(0, SCAN_BLOCK_SIZE) + scanned_rows = current_row - backward_offsets + valid_scanned_rows = scanned_rows >= 0 + scanned_req_idx = tl.load( + b_req_idx + scanned_rows, + mask=valid_scanned_rows, + other=-1, + ) + belongs_to_current_request = valid_scanned_rows & (scanned_req_idx == current_req_idx) + request_row_count = tl.sum(belongs_to_current_request, axis=0) + request_start_row = current_row - request_row_count + 1 + + # Split the request run from left to right into groups of MAX_GROUP_SIZE. + # Each loop iteration finds one group's final row and writes its size. + for group_start_offset in tl.range(0, request_row_count, MAX_GROUP_SIZE): + group_size = tl.minimum(MAX_GROUP_SIZE, request_row_count - group_start_offset) + group_end_row = request_start_row + group_start_offset + group_size - 1 + tl.store(b_mark_mtp_shared_group + group_end_row, group_size) + + +def build_mtp_shared_group_markers(b_req_idx: torch.Tensor, hold_req_id: int) -> torch.Tensor: + """Build MTP group markers from consecutive request indexes on GPU. + + Only the final row of each group stores its size. Long request runs are + split by ``max_batch_shared_group_size``. Every ``hold_req_id`` row forms + an independent one-row group. For example, with limit 3 and ``H`` denoting + ``hold_req_id``:: + + b_req_idx: [7, 7, 7, 7, 11, 11, H, H] + b_mark_mtp_shared_group: [0, 0, 3, 1, 0, 2, 1, 1] + """ + + assert b_req_idx.is_cuda + batch_size = b_req_idx.shape[0] + if batch_size == 0: + return torch.empty((0,), dtype=torch.int32, device=b_req_idx.device) + + max_group_size = int(get_diverse_max_batch_shared_group_size()) + assert max_group_size > 0 + b_mark_mtp_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device=b_req_idx.device) + _fwd_kernel_build_mtp_shared_group_markers[(batch_size,)]( + b_req_idx=b_req_idx, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, + batch_size=batch_size, + hold_req_id=hold_req_id, + MAX_GROUP_SIZE=max_group_size, + SCAN_BLOCK_SIZE=triton.next_power_of_2(batch_size), + num_warps=8, + num_stages=1, + ) + return b_mark_mtp_shared_group + @triton.jit def _fwd_kernel_mtp_verify( diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py index 1ebd95e3d3..ef78e3e96f 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py @@ -8,14 +8,12 @@ from lightllm.utils.infer_utils import calculate_time from lightllm.utils.envs_utils import ( enable_diverse_mode_gqa_decode_fast_kernel, - enable_triton_mtp_kernel, get_env_start_args, ) from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from .generic_pre_process import ( build_b_position_delta, build_diverse_shared_group_infos, - build_mtp_shared_group_markers, ) @@ -213,9 +211,6 @@ def padded_prepare_decode_inputs( if padded_row_count > 0: b_shared_seq_len = F.pad(b_shared_seq_len, (0, padded_row_count), value=0) b_mark_shared_group = F.pad(b_mark_shared_group, (0, padded_row_count), value=1) - elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): - b_shared_seq_len = None - b_mark_shared_group = build_mtp_shared_group_markers(b_req_idx=b_req_idx) else: b_shared_seq_len = None b_mark_shared_group = None diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 21dbb02cff..ae294544ce 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -5,9 +5,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.utils.envs_utils import ( enable_diverse_mode_gqa_decode_fast_kernel, - enable_triton_mtp_kernel, get_diverse_max_batch_shared_group_size, - get_env_start_args, ) @@ -136,9 +134,6 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In if enable_diverse_mode_gqa_decode_fast_kernel(): b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) - elif get_env_start_args().mtp_dynamic_verify or enable_triton_mtp_kernel(): - b_shared_seq_len = None - b_mark_shared_group = build_mtp_shared_group_markers(b_req_idx=b_req_idx) else: b_shared_seq_len = None b_mark_shared_group = None @@ -222,40 +217,3 @@ def build_diverse_shared_group_infos(run_reqs: List[InferReq]) -> Tuple[torch.Te b_shared_seq_len = torch.tensor(b_shared_seq_len, dtype=torch.int32, device="cpu") b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cpu") return b_shared_seq_len, b_mark_shared_group - - -def build_mtp_shared_group_markers(b_req_idx: torch.Tensor) -> torch.Tensor: - """Build MTP row-group markers from consecutive request indexes. - - Rows belonging to the same request form one MTP group. Only the final row - of a group stores the group size; all earlier rows store zero. A request - with more rows than the attention kernel supports is split into multiple - adjacent groups. - - Example with ``max_batch_shared_group_size == 3``:: - - b_req_idx: [7, 7, 7, 7, 11, 11, 20] - MTP groups: [7, 7, 7] [7] [11, 11] [20] - b_mark_shared_group: [0, 0, 3, 1, 0, 2, 1] - - The request id itself is not written to the result; it is only used to - detect where one request ends and the next request begins. - """ - max_batch_shared_group_size = get_diverse_max_batch_shared_group_size() - assert max_batch_shared_group_size > 0 - - req_indexes = b_req_idx.tolist() - b_mark_shared_group = [0] * len(req_indexes) - group_start = 0 - for row_index, req_idx in enumerate(req_indexes): - is_request_end = row_index == len(req_indexes) - 1 or req_indexes[row_index + 1] != req_idx - group_size = row_index - group_start + 1 - reaches_size_limit = group_size == max_batch_shared_group_size - if not is_request_end and not reaches_size_limit: - continue - - b_mark_shared_group[row_index] = group_size - group_start = row_index + 1 - - b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device=b_req_idx.device) - return b_mark_shared_group diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 2067855661..dbff6736b5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -156,7 +156,6 @@ def propose_next_dp_eagle_autoregressive_overlap( if position_deltas_by_batch[batch_index] is not None else None ) - model_input.b_mark_shared_group = torch.ones_like(model_input.b_req_idx) model_input.b_shared_seq_len = None if len(model_input.multimodal_params) != model_input.batch_size: empty_multimodal_params = {"images": [], "audios": []} diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 5c7b9e0005..902368c496 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -156,7 +156,6 @@ def propose_next_eagle( draft_input.b_position_delta = ( position_delta.index_select(0, accepted_tail_rows) if position_delta is not None else None ) - draft_input.b_mark_shared_group = torch.ones_like(draft_input.b_req_idx) draft_input.b_shared_seq_len = None if len(draft_input.multimodal_params) != request_count: empty_multimodal_params = {"images": [], "audios": []} diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py index 20d103d07a..ba661da2cb 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py @@ -93,8 +93,6 @@ def build_parallel_block_draft_input( .contiguous() ) draft_input.mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids.device, non_blocking=True) - draft_input.b_mark_shared_group = torch.zeros_like(draft_input.b_req_idx) - draft_input.b_mark_shared_group[block_size - 1 :: block_size] = block_size empty_multimodal_params = {"images": [], "audios": []} draft_input.multimodal_params = [empty_multimodal_params] * draft_input.batch_size return draft_input, extra_mem_indexes_cpu diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index caab276449..cb6ecfdc24 100644 --- a/unit_tests/common/basemodel/test_model_input.py +++ b/unit_tests/common/basemodel/test_model_input.py @@ -1,13 +1,10 @@ -from types import SimpleNamespace - -import pytest import torch import lightllm.common.basemodel.batch_objs as batch_objs_module from lightllm.common.basemodel.batch_objs import ModelInput -def _create_model_input(*, is_prefill=False, b_mtp_index=None): +def _create_model_input(*, is_prefill=False): batch_size = 2 return ModelInput( batch_size=batch_size, @@ -15,29 +12,23 @@ def _create_model_input(*, is_prefill=False, b_mtp_index=None): max_q_seq_len=1, max_kv_seq_len=1, b_req_idx=torch.arange(batch_size, dtype=torch.int32), - b_mtp_index=torch.zeros(batch_size, dtype=torch.int32) if b_mtp_index is None else b_mtp_index, + b_mtp_index=torch.zeros(batch_size, dtype=torch.int32), b_seq_len=torch.ones(batch_size, dtype=torch.int32), is_prefill=is_prefill, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], ) -def _mock_group_modes(monkeypatch, *, diverse=False, dynamic=False, triton=False): +def _mock_diverse_mode(monkeypatch, *, enabled=False): monkeypatch.setattr( batch_objs_module, "enable_diverse_mode_gqa_decode_fast_kernel", - lambda: diverse, - ) - monkeypatch.setattr(batch_objs_module, "enable_triton_mtp_kernel", lambda: triton) - monkeypatch.setattr( - batch_objs_module, - "get_env_start_args", - lambda: SimpleNamespace(mtp_dynamic_verify=dynamic), + lambda: enabled, ) def test_diverse_decode_defaults_to_independent_groups(monkeypatch): - _mock_group_modes(monkeypatch, diverse=True) + _mock_diverse_mode(monkeypatch, enabled=True) model_input = _create_model_input() model_input._ensure_decode_group_metadata() @@ -46,39 +37,18 @@ def test_diverse_decode_defaults_to_independent_groups(monkeypatch): assert torch.equal(model_input.b_shared_seq_len, torch.zeros(2, dtype=torch.int32)) -@pytest.mark.parametrize("dynamic,triton", [(True, False), (False, True)]) -def test_plain_spec_decode_defaults_to_single_row_groups(monkeypatch, dynamic, triton): - _mock_group_modes(monkeypatch, dynamic=dynamic, triton=triton) +def test_normal_decode_does_not_create_group_metadata(monkeypatch): + _mock_diverse_mode(monkeypatch) model_input = _create_model_input() model_input._ensure_decode_group_metadata() - assert torch.equal(model_input.b_mark_shared_group, torch.ones(2, dtype=torch.int32)) + assert model_input.b_mark_shared_group is None assert model_input.b_shared_seq_len is None -def test_multi_row_spec_decode_defaults_to_single_row_groups(monkeypatch): - _mock_group_modes(monkeypatch, dynamic=True) - model_input = _create_model_input(b_mtp_index=torch.tensor([0, 1], dtype=torch.int32)) - - model_input._ensure_decode_group_metadata() - - assert torch.equal(model_input.b_mark_shared_group, torch.ones(2, dtype=torch.int32)) - - -def test_multi_row_spec_decode_preserves_explicit_group_metadata(monkeypatch): - _mock_group_modes(monkeypatch, dynamic=True) - model_input = _create_model_input(b_mtp_index=torch.tensor([0, 1], dtype=torch.int32)) - group_markers = torch.tensor([0, 2], dtype=torch.int32) - model_input.b_mark_shared_group = group_markers - - model_input._ensure_decode_group_metadata() - - assert model_input.b_mark_shared_group is group_markers - - def test_prefill_does_not_create_decode_group_metadata(monkeypatch): - _mock_group_modes(monkeypatch, diverse=True, dynamic=True, triton=True) + _mock_diverse_mode(monkeypatch, enabled=True) model_input = _create_model_input(is_prefill=True) model_input._ensure_decode_group_metadata() diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py index 5c65d5e5a2..cc3a12a2ae 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py @@ -38,7 +38,6 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): ) monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "4") get_env_start_args.cache_clear() - dynamic_mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() model_input = ModelInput( batch_size=12, @@ -97,9 +96,7 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): assert torch.equal( compacted_input.b_shared_seq_len.cpu(), torch.tensor([0, 0, 0, 7, 9, 9, 9, 9], dtype=torch.int32) ) - assert torch.equal( - compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 1, 0, 0, 0, 4], dtype=torch.int32) - ) + assert compacted_input.b_mark_shared_group is None assert torch.equal( compacted_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 103, 104, 105, 106, 107], dtype=torch.int32) ) @@ -111,10 +108,7 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): assert torch.equal(compacted_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) -def test_compaction_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size(monkeypatch): - monkeypatch.setenv("LIGHTLLM_MAX_BATCH_SHARED_GROUP_SIZE", "3") - dynamic_mtp_utils.get_diverse_max_batch_shared_group_size.cache_clear() - +def test_compaction_clears_attention_group_metadata(): model_input = ModelInput( batch_size=5, total_token_num=36, @@ -137,13 +131,12 @@ def test_compaction_rebuilds_b_mark_shared_group_by_max_batch_shared_group_size( model_input=model_input, selected_row_mask=selected_row_mask, dynamic_batch_size=5, - max_draft_step=4, ) torch.cuda.synchronize() assert torch.equal(compacted_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32)) assert torch.equal(compacted_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32)) - assert torch.equal(compacted_input.b_mark_shared_group.cpu(), torch.tensor([0, 0, 3, 0, 2], dtype=torch.int32)) + assert compacted_input.b_mark_shared_group is None def _reference_cumprod_scores(req_to_next_token_scores, b_req_idx, max_draft_step: int) -> torch.Tensor: diff --git a/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py b/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py index 58255bcd36..c1ef686f8e 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_fa3_utils.py @@ -7,11 +7,11 @@ from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_spec_fa3_decode_params -def _reference_dynamic_spec_fa3_decode_params(b_req_idx, b_seq_len, b_mark_shared_group, hold_req_id): +def _reference_dynamic_spec_fa3_decode_params(b_req_idx, b_seq_len, b_mark_mtp_shared_group, hold_req_id): batch_size = b_req_idx.shape[0] b_req_idx_cpu = b_req_idx.cpu() b_seq_len_cpu = b_seq_len.cpu() - b_mark_cpu = b_mark_shared_group.cpu() + b_mark_cpu = b_mark_mtp_shared_group.cpu() valid_idx = torch.nonzero(b_mark_cpu > 0, as_tuple=False).flatten() valid_size = valid_idx.numel() @@ -32,24 +32,24 @@ def test_build_dynamic_spec_fa3_decode_params(batch_size): hold_req_id = -1 b_req_idx = torch.arange(batch_size, dtype=torch.int32, device="cuda") + 100 b_seq_len = (torch.arange(batch_size, dtype=torch.int32, device="cuda") % 97) + 1 - b_mark_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device="cuda") + b_mark_mtp_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device="cuda") mark_positions = [0, 3, 17, 255, 256, 511, batch_size - 1] for pos in mark_positions: if 0 <= pos < batch_size: - b_mark_shared_group[pos] = pos % 5 + 1 + b_mark_mtp_shared_group[pos] = pos % 5 + 1 actual = build_dynamic_spec_fa3_decode_params( b_req_idx=b_req_idx, b_seq_len=b_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, att_batch_size=batch_size, hold_req_id=hold_req_id, ) expected = _reference_dynamic_spec_fa3_decode_params( b_req_idx=b_req_idx, b_seq_len=b_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, hold_req_id=hold_req_id, ) @@ -62,19 +62,19 @@ def test_build_dynamic_spec_fa3_decode_params_all_padding(): hold_req_id = -1 b_req_idx = torch.arange(batch_size, dtype=torch.int32, device="cuda") b_seq_len = torch.arange(batch_size, dtype=torch.int32, device="cuda") + 1 - b_mark_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device="cuda") + b_mark_mtp_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device="cuda") actual = build_dynamic_spec_fa3_decode_params( b_req_idx=b_req_idx, b_seq_len=b_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, att_batch_size=batch_size, hold_req_id=hold_req_id, ) expected = _reference_dynamic_spec_fa3_decode_params( b_req_idx=b_req_idx, b_seq_len=b_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_mark_mtp_shared_group=b_mark_mtp_shared_group, hold_req_id=hold_req_id, ) diff --git a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py index b8206c78ad..f6f9df9d44 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -7,6 +7,24 @@ from lightllm.common.basemodel.triton_kernel import mtp_utils +@pytest.mark.parametrize( + "max_group_size,req_indexes,expected_markers", + [ + (2, [7, 7, 7, 11, 11, 20], [0, 2, 1, 0, 2, 1]), + (8, [7, 7, 7, 11, 11, 11, 20, 20, 20], [0, 0, 3, 0, 0, 3, 0, 0, 3]), + (3, [7, 7, 7, 7, 7, 7], [0, 0, 3, 0, 0, 3]), + (3, [7, 7, -1, -1, 11, 11], [0, 2, 1, 1, 0, 2]), + ], +) +def test_build_mtp_shared_group_markers(monkeypatch, max_group_size, req_indexes, expected_markers): + monkeypatch.setattr(mtp_utils, "get_diverse_max_batch_shared_group_size", lambda: max_group_size) + b_req_idx = torch.tensor(req_indexes, dtype=torch.int32, device="cuda") + + markers = mtp_utils.build_mtp_shared_group_markers(b_req_idx, hold_req_id=-1) + + assert markers.cpu().tolist() == expected_markers + + def test_mtp_verify_scatter_and_start_locations(): req_to_next_token_ids = torch.tensor( [[1, 2, -2, -1, -1], [1, 2, 0, -1, -1], [1, 3, 4, 4, 5]], diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index 8c8ff37d92..44876ccd6e 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -1,28 +1,10 @@ from types import SimpleNamespace import torch -from lightllm.server.router.model_infer.mode_backend import generic_padded_pre_process, generic_pre_process +from lightllm.server.router.model_infer.mode_backend import generic_padded_pre_process -def test_mtp_shared_group_markers_split_requests_and_size_limit(monkeypatch): - monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 2) - b_req_idx = torch.tensor([7, 7, 7, 11, 11, 20], dtype=torch.int32) - - markers = generic_pre_process.build_mtp_shared_group_markers(b_req_idx=b_req_idx) - - assert markers.tolist() == [0, 2, 1, 0, 2, 1] - - -def test_mtp_shared_group_markers_detect_request_boundaries(monkeypatch): - monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 8) - b_req_idx = torch.tensor([7, 7, 7, 11, 11, 11, 20, 20, 20], dtype=torch.int32) - - markers = generic_pre_process.build_mtp_shared_group_markers(b_req_idx=b_req_idx) - - assert markers.tolist() == [0, 0, 3, 0, 0, 3, 0, 0, 3] - - -def test_padded_decode_builds_spec_metadata_for_real_and_fake_rows(monkeypatch): +def test_padded_decode_leaves_mtp_attention_metadata_unset(monkeypatch): max_draft_step = 2 mem_manager = SimpleNamespace( HOLD_TOKEN_MEMINDEX=-1, @@ -40,8 +22,6 @@ def test_padded_decode_builds_spec_metadata_for_real_and_fake_rows(monkeypatch): lambda: SimpleNamespace(mtp_step=max_draft_step, mtp_dynamic_verify=True), ) monkeypatch.setattr(generic_padded_pre_process, "enable_diverse_mode_gqa_decode_fast_kernel", lambda: False) - monkeypatch.setattr(generic_padded_pre_process, "enable_triton_mtp_kernel", lambda: False) - monkeypatch.setattr(generic_pre_process, "get_diverse_max_batch_shared_group_size", lambda: 8) req = SimpleNamespace( req_idx=7, @@ -56,4 +36,4 @@ def test_padded_decode_builds_spec_metadata_for_real_and_fake_rows(monkeypatch): assert padded_req_num == 2 assert model_input.b_mtp_index.tolist() == [0, 1, 2, 0, 1, 2, 0, 1, 2] - assert model_input.b_mark_shared_group.tolist() == [0, 0, 3, 0, 0, 0, 0, 0, 6] + assert model_input.b_mark_shared_group is None diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index e17f946f4c..b743f98313 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -10,6 +10,7 @@ from lightllm.common.basemodel.attention.base_att import BaseAttBackend from lightllm.common.basemodel.attention.fa3.fp import Fa3DecodeAttState, Fa3PrefillAttState from lightllm.common.basemodel.attention.fa3.mla import MlaFa3DecodeAttState, MlaFa3PrefillAttState +from lightllm.common.basemodel.attention.triton.fp import TritonDecodeAttState from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight from lightllm.server.router.model_infer.mtp_speculative.proposers import dflash as dflash_module @@ -159,6 +160,46 @@ def test_fa3_decode_state_owns_causality(state_class): assert state.causal is False +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_triton_mtp_decode_state_builds_group_markers(): + model = SimpleNamespace( + is_mtp_draft_model=False, + mtp_manager=SimpleNamespace(get_decode_draft_step=lambda _: 2), + req_manager=SimpleNamespace(HOLD_REQUEST_ID=-1), + ) + state = TritonDecodeAttState( + backend=SimpleNamespace(model=model), + infer_state=SimpleNamespace( + b_req_idx=torch.tensor([7, 7, -1, -1], dtype=torch.int32, device="cuda"), + ), + ) + + state.init_state() + + assert state.b_mark_mtp_shared_group.cpu().tolist() == [0, 2, 1, 1] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("state_class", [Fa3DecodeAttState, MlaFa3DecodeAttState]) +def test_fa3_dynamic_decode_state_builds_group_markers(state_class): + model = SimpleNamespace(req_manager=SimpleNamespace(HOLD_REQUEST_ID=-1)) + state = state_class( + backend=SimpleNamespace(model=model), + infer_state=SimpleNamespace( + b_req_idx=torch.tensor([7, 7, -1, -1], dtype=torch.int32, device="cuda"), + b_seq_len=torch.tensor([3, 4, 2, 2], dtype=torch.int32, device="cuda"), + batch_size=4, + ), + ) + + b_att_req_idx = state._init_dynamic_spec_verify_state(draft_step=2) + torch.cuda.synchronize() + + assert b_att_req_idx.cpu().tolist() == [7, -1, -1, -1] + assert state.b_att_seq_len.cpu().tolist() == [4, 2, 2, 0] + assert state.cu_seqlens_q.cpu().tolist() == [0, 2, 3, 4, 4] + + @pytest.mark.parametrize( "model_type, spec_mode, expected_class_name", [ From ee69b52d6adbf417d40808ad081e205db0cd47dc Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 04:21:46 +0000 Subject: [PATCH 064/103] refactor: rebuild diverse groups from radix metadata --- .../basemodel/attention/triton/int8kv.py | 15 ++- lightllm/common/basemodel/basemodel.py | 20 ++-- lightllm/common/basemodel/batch_objs.py | 37 +++----- lightllm/common/basemodel/cuda_graph.py | 8 ++ lightllm/common/basemodel/infer_struct.py | 4 +- .../int8kv/int8kv_flash_decoding_diverse.py | 11 ++- .../basemodel/triton_kernel/diverse_utils.py | 91 +++++++++++++++++++ .../triton_kernel/dynamic_mtp_utils.py | 35 ++++--- .../generic_padded_pre_process.py | 33 ++++--- .../mode_backend/generic_pre_process.py | 71 ++++----------- .../dp_overlap_proposers/eagle_utils.py | 9 +- .../mtp_speculative/proposers/eagle_utils.py | 5 +- .../proposers/parallel_block_utils.py | 10 ++ .../common/basemodel/test_model_input.py | 61 ++++++------- .../test_int8kv_flash_decoding_diverse.py | 7 +- .../triton_kernel/test_diverse_utils.py | 46 ++++++++++ .../triton_kernel/test_dynamic_mtp_utils.py | 18 ++-- .../mode_backend/test_generic_pre_process.py | 9 +- .../mtp_speculative/test_eagle_overlap.py | 2 + unit_tests/utils/test_speculative_utils.py | 2 + 20 files changed, 319 insertions(+), 175 deletions(-) create mode 100644 lightllm/common/basemodel/triton_kernel/diverse_utils.py create mode 100644 unit_tests/common/basemodel/triton_kernel/test_diverse_utils.py diff --git a/lightllm/common/basemodel/attention/triton/int8kv.py b/lightllm/common/basemodel/attention/triton/int8kv.py index a47f63a40e..f2b3371ccd 100644 --- a/lightllm/common/basemodel/attention/triton/int8kv.py +++ b/lightllm/common/basemodel/attention/triton/int8kv.py @@ -4,6 +4,7 @@ from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl from typing import Optional, Tuple from lightllm.utils.envs_utils import enable_diverse_mode_gqa_decode_fast_kernel +from lightllm.common.basemodel.triton_kernel.diverse_utils import build_diverse_shared_group_markers class Int8kvTritonAttBackend(BaseAttBackend): @@ -115,8 +116,20 @@ def _groupsize_quant_prefill_att( @dataclasses.dataclass class Int8kvTritonDecodeAttState(BaseDecodeAttState): + b_shared_seq_len: torch.Tensor = None + b_mark_shared_group: torch.Tensor = None + def init_state(self): - pass + if enable_diverse_mode_gqa_decode_fast_kernel(): + self.b_mark_shared_group = build_diverse_shared_group_markers( + b_shared_radix_node_id=self.infer_state.b_shared_radix_node_id, + ) + # A one-row group has no cross-request prefix sharing to accelerate. + self.b_shared_seq_len = torch.where( + self.b_mark_shared_group == 1, + 0, + self.infer_state.b_shared_seq_len, + ) def copy_for_decode_cuda_graph(self, new_state: "Int8kvTritonDecodeAttState"): super().copy_for_decode_cuda_graph(new_state) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 0df46554e1..f9b9c9984c 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -35,7 +35,6 @@ from lightllm.utils.custom_kernel_utis import pad2dim_tensor_to_new_batch from lightllm.utils.envs_utils import ( set_model_init_status, - enable_diverse_mode_gqa_decode_fast_kernel, enable_full_att_decode_tune, ) from lightllm.common.triton_utils.autotuner import Autotuner @@ -400,9 +399,9 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) infer_state.b_ready_cache_len = model_input.b_ready_cache_len else: infer_state.b_ready_cache_len = torch.zeros_like(input=infer_state.b_seq_len) - elif enable_diverse_mode_gqa_decode_fast_kernel(): + else: infer_state.b_shared_seq_len = model_input.b_shared_seq_len - infer_state.b_mark_shared_group = model_input.b_mark_shared_group + infer_state.b_shared_radix_node_id = model_input.b_shared_radix_node_id infer_state.multimodal_params = model_input.multimodal_params @@ -464,15 +463,12 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s {"images": [], "audios": []} for _ in range(padded_batch_size) ] - if enable_diverse_mode_gqa_decode_fast_kernel(): - if new_model_input.b_shared_seq_len is not None: - new_model_input.b_shared_seq_len = F.pad( - new_model_input.b_shared_seq_len, (0, padded_batch_size), mode="constant", value=0 - ) - if new_model_input.b_mark_shared_group is not None: - new_model_input.b_mark_shared_group = F.pad( - new_model_input.b_mark_shared_group, (0, padded_batch_size), mode="constant", value=1 - ) + new_model_input.b_shared_seq_len = F.pad( + new_model_input.b_shared_seq_len, (0, padded_batch_size), mode="constant", value=0 + ) + new_model_input.b_shared_radix_node_id = F.pad( + new_model_input.b_shared_radix_node_id, (0, padded_batch_size), mode="constant", value=-1 + ) # 特殊模型,特殊模式的特殊变量的特殊 padding if new_model_input.mtp_draft_input_hiddens is not None: diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 68637874ef..04fe48aa25 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -4,7 +4,6 @@ import torch -from lightllm.utils.envs_utils import enable_diverse_mode_gqa_decode_fast_kernel from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor @@ -28,15 +27,12 @@ class ModelInput: b_is_decode_req: torch.Tensor = None - # 只会在 diverse_mode 下的 decode 阶段真正被使用的参数, 用于记录共享的radix cache中的长度 + # Decode 逐行携带的 radix cache 共享长度。普通 attention backend 不使用该 + # 信息,diverse attention backend 会结合 b_shared_radix_node_id 构建共享组。 b_shared_seq_len: torch.Tensor = None - # 只会在 diverse_mode 下的 decode 阶段真正被使用的参数, 用于记录请求间的共享关系。 - # 举列说明: - # b_shared_seq_len : [10, 10, 10, 11, 11, 11, 11] - # b_mark_shared_group: [0, 0, 3, 0, 0, 0, 4] - # b_mark_shared_group 中每一个不为0的位置都代表其与前面多少个请求形成一个共享前缀组。属于 - # 同一个共享前缀组的请求, 其在对应的 b_shared_seq_len 中的内容必然相同。 - b_mark_shared_group: torch.Tensor = None + # Decode 逐行携带的 radix node 标识。相同 id 表示请求引用同一个共享 + # radix node;该 id 只用于重建 diverse attention 的 b_mark_shared_group。 + b_shared_radix_node_id: torch.Tensor = None mem_indexes: torch.Tensor = None is_prefill: bool = False b_ready_cache_len: torch.Tensor = None @@ -85,11 +81,9 @@ def to_cuda(self): if self.b_prefill_start_loc is not None: self.b_prefill_start_loc = self.b_prefill_start_loc.cuda(non_blocking=True) - self._ensure_decode_group_metadata() - if self.b_mark_shared_group is not None: - self.b_mark_shared_group = self.b_mark_shared_group.cuda(non_blocking=True) - if self.b_shared_seq_len is not None: + if not self.is_prefill: self.b_shared_seq_len = self.b_shared_seq_len.cuda(non_blocking=True) + self.b_shared_radix_node_id = self.b_shared_radix_node_id.cuda(non_blocking=True) def __post_init__(self): self.check_input() @@ -100,18 +94,11 @@ def check_input(self): assert ( self.input_ids.dtype == torch.int64 ), f"model input_ids must use torch.int64, got {self.input_ids.dtype}" - - def _ensure_decode_group_metadata(self): - """按 decode 模式补齐能够安全降级的 attention 分组信息。""" - if self.is_prefill: - return - - if enable_diverse_mode_gqa_decode_fast_kernel(): - # 缺少共享信息时退化为互不共享前缀的单行组,不改变 attention 结果。 - if self.b_mark_shared_group is None: - self.b_mark_shared_group = torch.ones_like(self.b_req_idx, dtype=torch.int32) - if self.b_shared_seq_len is None: - self.b_shared_seq_len = torch.zeros_like(self.b_req_idx, dtype=torch.int32) + if not self.is_prefill: + assert self.b_shared_seq_len is not None + assert self.b_shared_radix_node_id is not None + assert self.b_shared_seq_len.shape == self.b_req_idx.shape + assert self.b_shared_radix_node_id.shape == self.b_req_idx.shape @dataclass diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index ad77f3aa70..5fde98eab8 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -258,6 +258,8 @@ def warmup(self, model): ) b_seq_len = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda") b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_shared_seq_len = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_shared_radix_node_id = torch.full((batch_size,), -1, dtype=torch.int64, device="cuda") model_input = ModelInput( batch_size=batch_size, @@ -269,6 +271,8 @@ def warmup(self, model): b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, + b_shared_seq_len=b_shared_seq_len, + b_shared_radix_node_id=b_shared_radix_node_id, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], @@ -315,6 +319,8 @@ def warmup_overlap(self, model): ) b_seq_len = torch.full((batch_size,), seq_len, dtype=torch.int32, device="cuda") b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_shared_seq_len = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_shared_radix_node_id = torch.full((batch_size,), -1, dtype=torch.int64, device="cuda") micro_batch = ModelInput( is_prefill=False, @@ -327,6 +333,8 @@ def warmup_overlap(self, model): mem_indexes=mem_indexes, b_req_idx=b_req_idx, b_seq_len=b_seq_len, + b_shared_seq_len=b_shared_seq_len, + b_shared_radix_node_id=b_shared_radix_node_id, b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cuda"), multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], **model._gen_special_model_input(batch_size), diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index d4835d2fc6..6434102dbe 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -34,8 +34,8 @@ def __init__(self): self.b_req_idx: torch.Tensor = None self.b_ready_cache_len: torch.Tensor = None # only for prefill prompt cache used. - self.b_shared_seq_len: torch.Tensor = None # only for diverse mode used in decode phase. - self.b_mark_shared_group: torch.Tensor = None # only for diverse mode used in decode phase. + self.b_shared_seq_len: torch.Tensor = None # raw decode radix-cache shared lengths. + self.b_shared_radix_node_id: torch.Tensor = None # raw decode radix-node ids. self.b_mtp_index: torch.Tensor = None # MRoPE position offset propagated from ModelInput. diff --git a/lightllm/common/basemodel/triton_kernel/att/decode_att/int8kv/int8kv_flash_decoding_diverse.py b/lightllm/common/basemodel/triton_kernel/att/decode_att/int8kv/int8kv_flash_decoding_diverse.py index b40251b8f8..c534795fdf 100644 --- a/lightllm/common/basemodel/triton_kernel/att/decode_att/int8kv/int8kv_flash_decoding_diverse.py +++ b/lightllm/common/basemodel/triton_kernel/att/decode_att/int8kv/int8kv_flash_decoding_diverse.py @@ -18,6 +18,9 @@ def token_decode_attention_flash_decoding( alloc_tensor_func=torch.empty, shared_streams_dict={}, ): + b_shared_seq_len = infer_state.decode_att_state.b_shared_seq_len + b_mark_shared_group = infer_state.decode_att_state.b_mark_shared_group + if "stream1" not in shared_streams_dict: shared_streams_dict["stream1"] = torch.cuda.Stream() if "stream2" not in shared_streams_dict: @@ -55,8 +58,8 @@ def token_decode_attention_flash_decoding( v_scale=cache_v_scale, Req_to_tokens=infer_state.req_manager.req_to_token_indexs, B_req_idx=infer_state.b_req_idx, - b_shared_seq_len=infer_state.b_shared_seq_len, - b_mark_shared_group=infer_state.b_mark_shared_group, + b_shared_seq_len=b_shared_seq_len, + b_mark_shared_group=b_mark_shared_group, max_len_in_batch=infer_state.max_kv_seq_len, mid_out=mid_o, mid_out_logsumexp=mid_o_logexpsum, @@ -74,7 +77,7 @@ def token_decode_attention_flash_decoding( Req_to_tokens=infer_state.req_manager.req_to_token_indexs, B_req_idx=infer_state.b_req_idx, B_Seqlen=infer_state.b_seq_len, - b_shared_seq_len=infer_state.b_shared_seq_len, + b_shared_seq_len=b_shared_seq_len, max_len_in_batch=infer_state.max_kv_seq_len, mid_out=mid_o, mid_out_logsumexp=mid_o_logexpsum, @@ -88,7 +91,7 @@ def token_decode_attention_flash_decoding( mid_out=mid_o, mid_out_logexpsum=mid_o_logexpsum, B_Seqlen=infer_state.b_seq_len, - b_shared_seq_len=infer_state.b_shared_seq_len, + b_shared_seq_len=b_shared_seq_len, O=o_tensor.view(calcu_shape1), block_seq=BLOCK_SEQ, ) diff --git a/lightllm/common/basemodel/triton_kernel/diverse_utils.py b/lightllm/common/basemodel/triton_kernel/diverse_utils.py new file mode 100644 index 0000000000..2bb9b89cb5 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/diverse_utils.py @@ -0,0 +1,91 @@ +import torch +import triton +import triton.language as tl + +from lightllm.utils.envs_utils import get_diverse_max_batch_shared_group_size + + +@triton.jit +def _fwd_kernel_build_diverse_shared_group_markers( + b_shared_radix_node_id, + b_mark_shared_group, + batch_size, + MAX_GROUP_SIZE: tl.constexpr, + SCAN_BLOCK_SIZE: tl.constexpr, +): + current_row = tl.program_id(axis=0) + current_node_id = tl.load(b_shared_radix_node_id + current_row) + + # -1 means that this row has no reusable radix node. This includes padded + # HOLD rows, and every such row must remain an independent group. + if current_node_id == -1: + tl.store(b_mark_shared_group + current_row, 1) + return + + has_next_row = current_row + 1 < batch_size + next_node_id = tl.load( + b_shared_radix_node_id + current_row + 1, + mask=has_next_row, + other=-1, + ) + is_shared_run_end = (~has_next_row) | (next_node_id != current_node_id) + if not is_shared_run_end: + return + + backward_offsets = tl.arange(0, SCAN_BLOCK_SIZE) + scanned_rows = current_row - backward_offsets + valid_scanned_rows = scanned_rows >= 0 + scanned_node_id = tl.load( + b_shared_radix_node_id + scanned_rows, + mask=valid_scanned_rows, + other=-1, + ) + has_same_node_id = valid_scanned_rows & (scanned_node_id == current_node_id) + mismatch_count = tl.cumsum((~has_same_node_id).to(tl.int32), axis=0) + belongs_to_shared_run = has_same_node_id & (mismatch_count == 0) + shared_run_size = tl.sum(belongs_to_shared_run, axis=0) + shared_run_start = current_row - shared_run_size + 1 + + for group_start_offset in tl.range(0, shared_run_size, MAX_GROUP_SIZE): + group_size = tl.minimum(MAX_GROUP_SIZE, shared_run_size - group_start_offset) + group_end_row = shared_run_start + group_start_offset + group_size - 1 + tl.store(b_mark_shared_group + group_end_row, group_size) + + +def build_diverse_shared_group_markers( + b_shared_radix_node_id: torch.Tensor, +) -> torch.Tensor: + """根据各 decode 行对应的 radix node ``time_id`` 构建 diverse 共享组标记。 + + 连续且相同的非负 ``time_id`` 会组成一个共享组;每个组只在最后一行写入 + 组大小,组内其他行写入 0。超过 ``max_group_size`` 的连续请求会被拆分成 + 多个共享组。``-1`` 表示该行没有可复用的 radix node,因此每个 ``-1`` + 都独立成组并写入 1。 + + 例如 ``max_group_size = 3`` 时:: + + b_shared_radix_node_id: [7, 7, 7, 7, 11, 11, -1, -1] + b_mark_shared_group: [0, 0, 3, 1, 0, 2, 1, 1] + + 前四个 ``7`` 被拆成大小为 3 和 1 的两个组,两个 ``11`` 组成大小为 2 + 的组,最后两个 ``-1`` 则分别作为独立的一行组。 + """ + + assert b_shared_radix_node_id.is_cuda + batch_size = b_shared_radix_node_id.shape[0] + if batch_size == 0: + return torch.empty((0,), dtype=torch.int32, device=b_shared_radix_node_id.device) + + max_group_size = int(get_diverse_max_batch_shared_group_size()) + assert max_group_size > 0 + b_mark_shared_group = torch.zeros((batch_size,), dtype=torch.int32, device=b_shared_radix_node_id.device) + _fwd_kernel_build_diverse_shared_group_markers[(batch_size,)]( + b_shared_radix_node_id=b_shared_radix_node_id, + b_mark_shared_group=b_mark_shared_group, + batch_size=batch_size, + MAX_GROUP_SIZE=max_group_size, + SCAN_BLOCK_SIZE=triton.next_power_of_2(batch_size), + num_warps=8, + num_stages=1, + ) + return b_mark_shared_group diff --git a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py index 0ca62dbd11..c93939fb8b 100644 --- a/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -98,12 +98,13 @@ def _fwd_kernel_compact_dynamic_mtp_model_input( out_b_position_delta, b_shared_seq_len, out_b_shared_seq_len, + b_shared_radix_node_id, + out_b_shared_radix_node_id, selected_mask, selected_dst_pos, batch_size, HAS_INPUT_IDS: tl.constexpr, HAS_B_POSITION_DELTA: tl.constexpr, - HAS_B_SHARED_SEQ_LEN: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): offsets = tl.arange(0, BLOCK_SIZE) @@ -130,9 +131,10 @@ def _fwd_kernel_compact_dynamic_mtp_model_input( position_delta = tl.load(b_position_delta + offsets, mask=mask, other=0) tl.store(out_b_position_delta + dst_pos, position_delta, mask=write_mask) - if HAS_B_SHARED_SEQ_LEN: - shared_seq_len = tl.load(b_shared_seq_len + offsets, mask=mask, other=0) - tl.store(out_b_shared_seq_len + dst_pos, shared_seq_len, mask=write_mask) + shared_seq_len = tl.load(b_shared_seq_len + offsets, mask=mask, other=0) + shared_radix_node_id = tl.load(b_shared_radix_node_id + offsets, mask=mask, other=-1) + tl.store(out_b_shared_seq_len + dst_pos, shared_seq_len, mask=write_mask) + tl.store(out_b_shared_radix_node_id + dst_pos, shared_radix_node_id, mask=write_mask) return @@ -210,6 +212,8 @@ def _compact_decode_model_input( assert model_input.b_req_idx.is_cuda assert model_input.b_mtp_index.is_cuda assert model_input.b_seq_len.is_cuda + assert model_input.b_shared_seq_len.is_cuda + assert model_input.b_shared_radix_node_id.is_cuda # Dynamic scheduling guarantees exactly dynamic_batch_size selected rows. selected_row_mask = selected_row_mask.to(torch.int32) @@ -223,12 +227,14 @@ def _compact_decode_model_input( (dynamic_batch_size,), dtype=model_input.input_ids.dtype, device=model_input.input_ids.device ) - out_b_shared_seq_len = None - if model_input.b_shared_seq_len is not None: - assert model_input.b_shared_seq_len.is_cuda - out_b_shared_seq_len = torch.empty( - (dynamic_batch_size,), dtype=model_input.b_shared_seq_len.dtype, device=model_input.b_shared_seq_len.device - ) + out_b_shared_seq_len = torch.empty( + (dynamic_batch_size,), dtype=model_input.b_shared_seq_len.dtype, device=model_input.b_shared_seq_len.device + ) + out_b_shared_radix_node_id = torch.empty( + (dynamic_batch_size,), + dtype=model_input.b_shared_radix_node_id.dtype, + device=model_input.b_shared_radix_node_id.device, + ) out_b_req_idx = torch.empty( (dynamic_batch_size,), dtype=model_input.b_req_idx.dtype, device=model_input.b_req_idx.device @@ -262,14 +268,15 @@ def _compact_decode_model_input( out_b_seq_len=out_b_seq_len, b_position_delta=model_input.b_position_delta if model_input.b_position_delta is not None else dummy_1d, out_b_position_delta=out_b_position_delta if out_b_position_delta is not None else dummy_1d, - b_shared_seq_len=model_input.b_shared_seq_len if model_input.b_shared_seq_len is not None else dummy_1d, - out_b_shared_seq_len=out_b_shared_seq_len if out_b_shared_seq_len is not None else dummy_1d, + b_shared_seq_len=model_input.b_shared_seq_len, + out_b_shared_seq_len=out_b_shared_seq_len, + b_shared_radix_node_id=model_input.b_shared_radix_node_id, + out_b_shared_radix_node_id=out_b_shared_radix_node_id, selected_mask=selected_row_mask, selected_dst_pos=selected_dst_pos, batch_size=old_batch_size, HAS_INPUT_IDS=model_input.input_ids is not None, HAS_B_POSITION_DELTA=model_input.b_position_delta is not None, - HAS_B_SHARED_SEQ_LEN=model_input.b_shared_seq_len is not None, BLOCK_SIZE=BLOCK_SIZE, num_warps=8, num_stages=1, @@ -281,7 +288,7 @@ def _compact_decode_model_input( model_input.b_seq_len = out_b_seq_len model_input.b_position_delta = out_b_position_delta model_input.b_shared_seq_len = out_b_shared_seq_len - model_input.b_mark_shared_group = None + model_input.b_shared_radix_node_id = out_b_shared_radix_node_id if model_input.mtp_draft_input_hiddens is not None: assert model_input.mtp_draft_input_hiddens.is_cuda diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py index ef78e3e96f..d829163fb1 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py @@ -6,14 +6,11 @@ from typing import List, Optional, Tuple from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq from lightllm.utils.infer_utils import calculate_time -from lightllm.utils.envs_utils import ( - enable_diverse_mode_gqa_decode_fast_kernel, - get_env_start_args, -) +from lightllm.utils.envs_utils import get_env_start_args from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from .generic_pre_process import ( + INT64_MAX, build_b_position_delta, - build_diverse_shared_group_infos, ) @@ -206,14 +203,22 @@ def padded_prepare_decode_inputs( b_position_delta = build_b_position_delta(batch_multimodal_params) padded_row_count = padded_req_num * (args_mtp_step + 1) - if enable_diverse_mode_gqa_decode_fast_kernel(): - b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) - if padded_row_count > 0: - b_shared_seq_len = F.pad(b_shared_seq_len, (0, padded_row_count), value=0) - b_mark_shared_group = F.pad(b_mark_shared_group, (0, padded_row_count), value=1) - else: - b_shared_seq_len = None - b_mark_shared_group = None + b_shared_seq_len = torch.tensor( + [req.get_radix_cache_shared_len() for req in run_reqs], dtype=torch.int32, device="cpu" + ) + b_shared_radix_node_id = torch.tensor( + [ + -1 + if req.shared_kv_node is None + else req.shared_kv_node.time_id % INT64_MAX + for req in run_reqs + ], + dtype=torch.int64, + device="cpu", + ) + if padded_row_count > 0: + b_shared_seq_len = F.pad(b_shared_seq_len, (0, padded_row_count), value=0) + b_shared_radix_node_id = F.pad(b_shared_radix_node_id, (0, padded_row_count), value=-1) # dynamic prompt cache 准备 token if g_infer_context.radix_cache is not None: @@ -240,7 +245,7 @@ def padded_prepare_decode_inputs( b_seq_len=b_seq_len, b_position_delta=b_position_delta, b_shared_seq_len=b_shared_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_shared_radix_node_id=b_shared_radix_node_id, is_prefill=False, multimodal_params=batch_multimodal_params, ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index ae294544ce..6d9920b585 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -3,10 +3,8 @@ from typing import List, Tuple from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context from lightllm.common.basemodel.batch_objs import ModelInput -from lightllm.utils.envs_utils import ( - enable_diverse_mode_gqa_decode_fast_kernel, - get_diverse_max_batch_shared_group_size, -) + +INT64_MAX = torch.iinfo(torch.int64).max def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> Tuple[ModelInput, List[InferReq]]: @@ -132,11 +130,19 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_mtp_index = torch.tensor(b_mtp_index, dtype=torch.int32, device="cpu") b_position_delta = build_b_position_delta(multimodal_params) - if enable_diverse_mode_gqa_decode_fast_kernel(): - b_shared_seq_len, b_mark_shared_group = build_diverse_shared_group_infos(run_reqs=run_reqs) - else: - b_shared_seq_len = None - b_mark_shared_group = None + b_shared_seq_len = torch.tensor( + [req.get_radix_cache_shared_len() for req in run_reqs], dtype=torch.int32, device="cpu" + ) + b_shared_radix_node_id = torch.tensor( + [ + -1 + if req.shared_kv_node is None + else req.shared_kv_node.time_id % INT64_MAX + for req in run_reqs + ], + dtype=torch.int64, + device="cpu", + ) # dynamic prompt cache 准备 token if g_infer_context.radix_cache is not None: @@ -155,7 +161,7 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In b_seq_len=b_seq_len, b_position_delta=b_position_delta, b_shared_seq_len=b_shared_seq_len, - b_mark_shared_group=b_mark_shared_group, + b_shared_radix_node_id=b_shared_radix_node_id, is_prefill=False, multimodal_params=multimodal_params, ) @@ -172,48 +178,3 @@ def build_b_position_delta(multimodal_params: List[dict]) -> torch.Tensor: position_delta += grid_thwd[3] b_position_delta.append(position_delta) return torch.tensor(b_position_delta, dtype=torch.int32, device="cpu") - - -def build_diverse_shared_group_infos(run_reqs: List[InferReq]) -> Tuple[torch.Tensor, torch.Tensor]: - # b_shared_seq_len 和 b_mark_shared_group 只会在 diverse_mode 下的 decode 阶段真正被使用的参数, - # 用于记录请求间的共享关系。 - # 举列说明: - # b_shared_seq_len : [10, 10, 10, 11, 11, 11, 11] - # b_mark_shared_group: [0, 0, 3, 0, 0, 0, 4] - # b_mark_shared_group 中每一个不为0的位置都代表其与前面多少个请求形成一个共享前缀组。属于 - # 同一个共享前缀组的请求, 其在对应的 b_shared_seq_len 中的内容必然相同。某些模式可以利用这两个 - # 输入加速算子的运行。 - max_batch_shared_group_size = get_diverse_max_batch_shared_group_size() - b_shared_seq_len = [req.get_radix_cache_shared_len() for req in run_reqs] - b_mark_shared_group = [] - shared_nodes = [req.shared_kv_node for req in run_reqs] - _current_group = [] - for node in shared_nodes: - if not _current_group: - _current_group.append(node) - elif node == _current_group[-1]: - _current_group.append(node) - else: - b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) - b_mark_shared_group[-1] = len(_current_group) - _current_group.clear() - _current_group.append(node) - - if len(_current_group) == max_batch_shared_group_size: - b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) - b_mark_shared_group[-1] = len(_current_group) - _current_group.clear() - if _current_group: - b_mark_shared_group.extend([0 for _ in range(len(_current_group))]) - b_mark_shared_group[-1] = len(_current_group) - _current_group.clear() - - assert len(b_mark_shared_group) == len(run_reqs) - # 如果一个 shared group 的长度为1, 则将其共享长度强制修改为0, 避免无效计算,提升 - # 算子执行效率。 - b_shared_seq_len = [ - 0 if group_size == 1 else shared_len for shared_len, group_size in zip(b_shared_seq_len, b_mark_shared_group) - ] - b_shared_seq_len = torch.tensor(b_shared_seq_len, dtype=torch.int32, device="cpu") - b_mark_shared_group = torch.tensor(b_mark_shared_group, dtype=torch.int32, device="cpu") - return b_shared_seq_len, b_mark_shared_group diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index dbff6736b5..59e1289976 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -123,6 +123,8 @@ def propose_next_dp_eagle_autoregressive_overlap( draft_hiddens_by_batch = [] draft_seq_lens_by_batch = [] draft_req_indices_by_batch = [] + draft_shared_seq_lens_by_batch = [] + draft_shared_radix_node_ids_by_batch = [] proposal_row_offsets = (0, real_request_counts[0]) for batch_index, (model_input, extend_output, accepted_tail_rows, real_request_count) in enumerate( zip(model_inputs, extend_outputs, accepted_tail_rows_by_batch, real_request_counts) @@ -133,6 +135,10 @@ def propose_next_dp_eagle_autoregressive_overlap( draft_hiddens_by_batch.append(extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows)) draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) + draft_shared_seq_lens_by_batch.append(model_input.b_shared_seq_len.index_select(0, accepted_tail_rows)) + draft_shared_radix_node_ids_by_batch.append( + model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) + ) proposal_row_start = proposal_row_offsets[batch_index] proposal_token_ids[proposal_row_start : proposal_row_start + real_request_count, 0] = draft_token_ids[ :real_request_count @@ -156,7 +162,8 @@ def propose_next_dp_eagle_autoregressive_overlap( if position_deltas_by_batch[batch_index] is not None else None ) - model_input.b_shared_seq_len = None + model_input.b_shared_seq_len = draft_shared_seq_lens_by_batch[batch_index] + model_input.b_shared_radix_node_id = draft_shared_radix_node_ids_by_batch[batch_index] if len(model_input.multimodal_params) != model_input.batch_size: empty_multimodal_params = {"images": [], "audios": []} model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 902368c496..4997236faa 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -156,7 +156,10 @@ def propose_next_eagle( draft_input.b_position_delta = ( position_delta.index_select(0, accepted_tail_rows) if position_delta is not None else None ) - draft_input.b_shared_seq_len = None + draft_input.b_shared_seq_len = target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) + draft_input.b_shared_radix_node_id = target_model_input.b_shared_radix_node_id.index_select( + 0, accepted_tail_rows + ) if len(draft_input.multimodal_params) != request_count: empty_multimodal_params = {"images": [], "audios": []} draft_input.multimodal_params = [empty_multimodal_params] * request_count diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py index ba661da2cb..c8824a8e23 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py @@ -92,6 +92,16 @@ def build_parallel_block_draft_input( .repeat_interleave(block_size) .contiguous() ) + draft_input.b_shared_seq_len = ( + target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() + ) + draft_input.b_shared_radix_node_id = ( + target_model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() + ) draft_input.mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids.device, non_blocking=True) empty_multimodal_params = {"images": [], "audios": []} draft_input.multimodal_params = [empty_multimodal_params] * draft_input.batch_size diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index cb6ecfdc24..30de454283 100644 --- a/unit_tests/common/basemodel/test_model_input.py +++ b/unit_tests/common/basemodel/test_model_input.py @@ -1,12 +1,12 @@ +import pytest import torch -import lightllm.common.basemodel.batch_objs as batch_objs_module from lightllm.common.basemodel.batch_objs import ModelInput def _create_model_input(*, is_prefill=False): batch_size = 2 - return ModelInput( + kwargs = dict( batch_size=batch_size, total_token_num=batch_size, max_q_seq_len=1, @@ -17,41 +17,36 @@ def _create_model_input(*, is_prefill=False): is_prefill=is_prefill, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], ) - - -def _mock_diverse_mode(monkeypatch, *, enabled=False): - monkeypatch.setattr( - batch_objs_module, - "enable_diverse_mode_gqa_decode_fast_kernel", - lambda: enabled, - ) - - -def test_diverse_decode_defaults_to_independent_groups(monkeypatch): - _mock_diverse_mode(monkeypatch, enabled=True) + if not is_prefill: + kwargs["b_shared_seq_len"] = torch.tensor([4, 4], dtype=torch.int32) + kwargs["b_shared_radix_node_id"] = torch.tensor([10, 10], dtype=torch.int64) + return ModelInput(**kwargs) + + +def test_decode_requires_shared_radix_metadata(): + with pytest.raises(AssertionError): + ModelInput( + batch_size=1, + total_token_num=1, + max_q_seq_len=1, + max_kv_seq_len=1, + b_req_idx=torch.zeros(1, dtype=torch.int32), + b_mtp_index=torch.zeros(1, dtype=torch.int32), + b_seq_len=torch.ones(1, dtype=torch.int32), + is_prefill=False, + multimodal_params=[{"images": [], "audios": []}], + ) + + +def test_decode_carries_raw_shared_radix_metadata(): model_input = _create_model_input() - model_input._ensure_decode_group_metadata() - - assert torch.equal(model_input.b_mark_shared_group, torch.ones(2, dtype=torch.int32)) - assert torch.equal(model_input.b_shared_seq_len, torch.zeros(2, dtype=torch.int32)) + assert torch.equal(model_input.b_shared_seq_len, torch.tensor([4, 4], dtype=torch.int32)) + assert torch.equal(model_input.b_shared_radix_node_id, torch.tensor([10, 10], dtype=torch.int64)) -def test_normal_decode_does_not_create_group_metadata(monkeypatch): - _mock_diverse_mode(monkeypatch) - model_input = _create_model_input() - - model_input._ensure_decode_group_metadata() - - assert model_input.b_mark_shared_group is None - assert model_input.b_shared_seq_len is None - - -def test_prefill_does_not_create_decode_group_metadata(monkeypatch): - _mock_diverse_mode(monkeypatch, enabled=True) +def test_prefill_does_not_require_shared_radix_metadata(): model_input = _create_model_input(is_prefill=True) - model_input._ensure_decode_group_metadata() - - assert model_input.b_mark_shared_group is None assert model_input.b_shared_seq_len is None + assert model_input.b_shared_radix_node_id is None diff --git a/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py b/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py index 3e2555b339..ec70d3c697 100644 --- a/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py +++ b/unit_tests/common/basemodel/triton_kernel/att/decode_att/int8kv/test_int8kv_flash_decoding_diverse.py @@ -1,4 +1,5 @@ import pytest +from types import SimpleNamespace import torch @@ -25,16 +26,12 @@ def __init__( req_to_tokens, b_req_idx, b_seq_len, - b_shared_seq_len=None, - b_mark_shared_group=None, ): self.batch_size = batch_size self.max_kv_seq_len = max_kv_seq_len self.req_manager = MockReqManager(req_to_tokens) self.b_req_idx = b_req_idx self.b_seq_len = b_seq_len - self.b_shared_seq_len = b_shared_seq_len - self.b_mark_shared_group = b_mark_shared_group # @pytest.mark.parametrize("shared_seq_len", [512]) @@ -104,6 +101,8 @@ def test_token_decode_attention_flash_decoding_diverse_matches_normal_decode(sha req_to_tokens=req_to_tokens, b_req_idx=b_req_idx, b_seq_len=b_seq_len, + ) + diverse_infer_state.decode_att_state = SimpleNamespace( b_shared_seq_len=b_shared_seq_len, b_mark_shared_group=b_mark_shared_group, ) diff --git a/unit_tests/common/basemodel/triton_kernel/test_diverse_utils.py b/unit_tests/common/basemodel/triton_kernel/test_diverse_utils.py new file mode 100644 index 0000000000..357975fbb3 --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_diverse_utils.py @@ -0,0 +1,46 @@ +from types import SimpleNamespace + +import pytest +import torch + +if not torch.cuda.is_available(): + pytest.skip("requires CUDA", allow_module_level=True) + +from lightllm.common.basemodel.attention.triton import int8kv as int8kv_module +from lightllm.common.basemodel.attention.triton.int8kv import Int8kvTritonDecodeAttState +from lightllm.common.basemodel.triton_kernel import diverse_utils + + +@pytest.mark.parametrize( + "node_ids,expected_markers", + [ + ([10, 10, 10, 10, 20, 20], [0, 0, 3, 1, 0, 2]), + ([10, 10, 20, 10, 10], [0, 2, 1, 0, 2]), + ([10, 10, -1, -1], [0, 2, 1, 1]), + ([-1, -1, -1, -1], [1, 1, 1, 1]), + ], +) +def test_build_diverse_shared_group_markers(monkeypatch, node_ids, expected_markers): + monkeypatch.setattr(diverse_utils, "get_diverse_max_batch_shared_group_size", lambda: 3) + b_shared_radix_node_id = torch.tensor(node_ids, dtype=torch.int64, device="cuda") + + markers = diverse_utils.build_diverse_shared_group_markers( + b_shared_radix_node_id=b_shared_radix_node_id, + ) + + assert markers.cpu().tolist() == expected_markers + + +def test_int8kv_decode_state_rebuilds_diverse_metadata(monkeypatch): + monkeypatch.setattr(int8kv_module, "enable_diverse_mode_gqa_decode_fast_kernel", lambda: True) + monkeypatch.setattr(diverse_utils, "get_diverse_max_batch_shared_group_size", lambda: 3) + infer_state = SimpleNamespace( + b_shared_seq_len=torch.tensor([8, 8, 5, 0], dtype=torch.int32, device="cuda"), + b_shared_radix_node_id=torch.tensor([10, 10, 20, -1], dtype=torch.int64, device="cuda"), + ) + state = Int8kvTritonDecodeAttState(backend=SimpleNamespace(), infer_state=infer_state) + + state.init_state() + + assert state.b_mark_shared_group.cpu().tolist() == [0, 2, 1, 1] + assert state.b_shared_seq_len.cpu().tolist() == [8, 8, 0, 0] diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py index cc3a12a2ae..42c4a44f68 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py @@ -50,7 +50,9 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): b_seq_len=torch.tensor([3, 4, 5, 6, 3, 4, 5, 6, 3, 4, 5, 6], dtype=torch.int32, device="cuda"), b_position_delta=torch.arange(12, dtype=torch.int32, device="cuda") + 200, b_shared_seq_len=torch.tensor([0, 0, 0, 0, 7, 7, 7, 7, 9, 9, 9, 9], dtype=torch.int32, device="cuda"), - b_mark_shared_group=torch.tensor([0, 0, 0, 4, 0, 0, 0, 4, 0, 0, 0, 4], dtype=torch.int32, device="cuda"), + b_shared_radix_node_id=torch.tensor( + [10, 10, 10, 10, 11, 11, 11, 11, 12, 12, 12, 12], dtype=torch.int64, device="cuda" + ), mem_indexes=torch.arange(8, dtype=torch.int32, device="cuda") + 100, mem_indexes_cpu=torch.arange(8, dtype=torch.int32, device="cpu") + 100, is_prefill=False, @@ -96,7 +98,10 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): assert torch.equal( compacted_input.b_shared_seq_len.cpu(), torch.tensor([0, 0, 0, 7, 9, 9, 9, 9], dtype=torch.int32) ) - assert compacted_input.b_mark_shared_group is None + assert torch.equal( + compacted_input.b_shared_radix_node_id.cpu(), + torch.tensor([10, 10, 10, 11, 12, 12, 12, 12], dtype=torch.int64), + ) assert torch.equal( compacted_input.mem_indexes.cpu(), torch.tensor([100, 101, 102, 103, 104, 105, 106, 107], dtype=torch.int32) ) @@ -108,7 +113,7 @@ def test_compact_dynamic_mtp_model_input(monkeypatch): assert torch.equal(compacted_input.mtp_draft_input_hiddens.cpu(), expected_hiddens) -def test_compaction_clears_attention_group_metadata(): +def test_compaction_preserves_shared_radix_metadata(): model_input = ModelInput( batch_size=5, total_token_num=36, @@ -118,8 +123,8 @@ def test_compaction_clears_attention_group_metadata(): b_req_idx=torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32, device="cuda"), b_mtp_index=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device="cuda"), b_seq_len=torch.tensor([3, 4, 5, 6, 7], dtype=torch.int32, device="cuda"), - b_shared_seq_len=None, - b_mark_shared_group=torch.tensor([0, 0, 0, 0, 5], dtype=torch.int32, device="cuda"), + b_shared_seq_len=torch.full((5,), 7, dtype=torch.int32, device="cuda"), + b_shared_radix_node_id=torch.full((5,), 10, dtype=torch.int64, device="cuda"), mem_indexes=torch.arange(5, dtype=torch.int32, device="cuda"), mem_indexes_cpu=torch.arange(5, dtype=torch.int32, device="cpu"), is_prefill=False, @@ -136,7 +141,8 @@ def test_compaction_clears_attention_group_metadata(): assert torch.equal(compacted_input.b_req_idx.cpu(), torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32)) assert torch.equal(compacted_input.b_mtp_index.cpu(), torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32)) - assert compacted_input.b_mark_shared_group is None + assert torch.equal(compacted_input.b_shared_seq_len.cpu(), torch.full((5,), 7, dtype=torch.int32)) + assert torch.equal(compacted_input.b_shared_radix_node_id.cpu(), torch.full((5,), 10, dtype=torch.int64)) def _reference_cumprod_scores(req_to_next_token_scores, b_req_idx, max_draft_step: int) -> torch.Tensor: diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index 44876ccd6e..6ed1324ec6 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -4,7 +4,7 @@ from lightllm.server.router.model_infer.mode_backend import generic_padded_pre_process -def test_padded_decode_leaves_mtp_attention_metadata_unset(monkeypatch): +def test_padded_decode_builds_raw_shared_radix_metadata(monkeypatch): max_draft_step = 2 mem_manager = SimpleNamespace( HOLD_TOKEN_MEMINDEX=-1, @@ -21,14 +21,16 @@ def test_padded_decode_leaves_mtp_attention_metadata_unset(monkeypatch): "get_env_start_args", lambda: SimpleNamespace(mtp_step=max_draft_step, mtp_dynamic_verify=True), ) - monkeypatch.setattr(generic_padded_pre_process, "enable_diverse_mode_gqa_decode_fast_kernel", lambda: False) + shared_kv_node = SimpleNamespace(time_id=torch.iinfo(torch.int64).max + 42) req = SimpleNamespace( req_idx=7, cur_kv_len=4, mtp_step=max_draft_step, multimodal_params={"images": [], "audios": []}, + shared_kv_node=shared_kv_node, get_cur_total_len=lambda: 5, + get_radix_cache_shared_len=lambda: 4, ) model_input, _, padded_req_num = generic_padded_pre_process.padded_prepare_decode_inputs( req_objs=[req], dest_batch_size=3 @@ -36,4 +38,5 @@ def test_padded_decode_leaves_mtp_attention_metadata_unset(monkeypatch): assert padded_req_num == 2 assert model_input.b_mtp_index.tolist() == [0, 1, 2, 0, 1, 2, 0, 1, 2] - assert model_input.b_mark_shared_group is None + assert model_input.b_shared_seq_len.tolist() == [4, 4, 4, 0, 0, 0, 0, 0, 0] + assert model_input.b_shared_radix_node_id.tolist() == [42, 42, 42, -1, -1, -1, -1, -1, -1] diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 40f11d218a..c8a59c58c1 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -54,6 +54,8 @@ def _target_input(batch_size): b_req_idx=torch.arange(batch_size, dtype=torch.int32), b_mtp_index=torch.zeros(batch_size, dtype=torch.int32), b_position_delta=torch.zeros(batch_size, dtype=torch.int32), + b_shared_seq_len=torch.zeros(batch_size, dtype=torch.int32), + b_shared_radix_node_id=torch.arange(batch_size, dtype=torch.int64), mem_indexes=torch.arange(batch_size, dtype=torch.int32), mem_indexes_cpu=torch.arange(batch_size, dtype=torch.int32), max_kv_seq_len=16, diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index b743f98313..711e9505a8 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -364,6 +364,8 @@ def test_dflash_expands_position_delta_with_request_block_rows(monkeypatch): b_req_idx=torch.tensor([10, 10, 11, 12, 12], dtype=torch.int32), b_seq_len=torch.tensor([4, 5, 7, 8, 9], dtype=torch.int32), b_position_delta=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32), + b_shared_seq_len=torch.tensor([3, 3, 0, 6, 6], dtype=torch.int32), + b_shared_radix_node_id=torch.tensor([1, 1, 2, 3, 3], dtype=torch.int64), max_kv_seq_len=9, ) From b3bf3198b5d41fae308f19ca8a10ae23241a471b Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 06:48:55 +0000 Subject: [PATCH 065/103] refactor: tighten model input phase contracts --- lightllm/common/basemodel/basemodel.py | 13 ++- lightllm/common/basemodel/batch_objs.py | 88 +++++++++------ .../common/basemodel/prefill_cuda_graph.py | 4 + lightllm/models/qwen2_vl/infer_struct.py | 32 +----- .../mode_backend/mtp_pre_process.py | 5 +- .../dp_overlap_proposers/eagle_utils.py | 13 +-- .../mtp_speculative/proposers/eagle_utils.py | 26 ++--- .../proposers/parallel_block_utils.py | 15 +-- .../proposers/vanilla_no_att.py | 75 +++++++++++-- .../common/basemodel/test_model_input.py | 84 +++++++++++++- .../triton_kernel/test_dynamic_mtp_utils.py | 1 + .../models/qwen2_vl/test_infer_struct.py | 31 +---- .../mtp_speculative/test_eagle_overlap.py | 9 +- .../mtp_speculative/test_vanilla_no_att.py | 106 ++++++++++++++++++ unit_tests/utils/test_speculative_utils.py | 15 ++- 15 files changed, 364 insertions(+), 153 deletions(-) create mode 100644 unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index f9b9c9984c..ce35c4158e 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -507,6 +507,7 @@ def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle new_model_input.b_mtp_index = F.pad(new_model_input.b_mtp_index, (0, 1), mode="constant", value=0) new_model_input.b_seq_len = F.pad(new_model_input.b_seq_len, (0, 1), mode="constant", value=padded_token_num) new_model_input.b_ready_cache_len = F.pad(new_model_input.b_ready_cache_len, (0, 1), mode="constant", value=0) + new_model_input.b_is_decode_req = F.pad(new_model_input.b_is_decode_req, (0, 1), mode="constant", value=False) b_q_seq_len = new_model_input.b_seq_len - new_model_input.b_ready_cache_len new_model_input.b_prefill_start_loc = b_q_seq_len.cumsum(dim=0, dtype=torch.int32) - b_q_seq_len # 构建新的list, 使用 append 可能会让外面使用的数组引用发生变化,导致错误。 @@ -559,7 +560,7 @@ def _prefill( self, model_input: ModelInput, ): - if self.args.enable_prefill_decode_mixed and model_input.b_is_decode_req is not None: + if self.args.enable_prefill_decode_mixed: gather_token_prefill_decode_mixed( input_ids=model_input.input_ids, req_to_next_token_ids=self.req_manager.req_sampling_params_manager.req_to_next_token_ids, @@ -776,7 +777,7 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod model_input0.to_cuda() model_input1.to_cuda() - if self.args.enable_prefill_decode_mixed and model_input0.b_is_decode_req is not None: + if self.args.enable_prefill_decode_mixed: gather_token_prefill_decode_mixed( input_ids=model_input0.input_ids, req_to_next_token_ids=self.req_manager.req_sampling_params_manager.req_to_next_token_ids, @@ -786,7 +787,7 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod b_prefill_start_loc=model_input0.b_prefill_start_loc, ) - if self.args.enable_prefill_decode_mixed and model_input1.b_is_decode_req is not None: + if self.args.enable_prefill_decode_mixed: gather_token_prefill_decode_mixed( input_ids=model_input1.input_ids, req_to_next_token_ids=self.req_manager.req_sampling_params_manager.req_to_next_token_ids, @@ -1100,6 +1101,7 @@ def _check_max_len_infer(self): b_prefill_start_loc = torch.zeros(1, dtype=torch.int32, device="cuda") total_token_num = self.batch_max_tokens b_mtp_index = torch.zeros(1, dtype=torch.int32, device="cuda") + b_is_decode_req = torch.zeros(1, dtype=torch.bool, device="cuda") model_input = ModelInput( batch_size=1, total_token_num=total_token_num, @@ -1112,6 +1114,7 @@ def _check_max_len_infer(self): b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, + b_is_decode_req=b_is_decode_req, is_prefill=True, b_ready_cache_len=b_ready_cache_len, b_prefill_start_loc=b_prefill_start_loc, @@ -1177,6 +1180,7 @@ def _autotune_warmup(self): b_prefill_start_loc = torch.zeros(1, dtype=torch.int32, device="cuda") total_token_num = input_len b_mtp_index = torch.zeros(1, dtype=torch.int32, device="cuda") + b_is_decode_req = torch.zeros(1, dtype=torch.bool, device="cuda") model_input = ModelInput( batch_size=1, total_token_num=total_token_num, @@ -1189,6 +1193,7 @@ def _autotune_warmup(self): b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, + b_is_decode_req=b_is_decode_req, is_prefill=True, b_ready_cache_len=b_ready_cache_len, b_prefill_start_loc=b_prefill_start_loc, @@ -1242,6 +1247,7 @@ def _init_padded_req(self): b_prefill_start_loc = b_q_seq_len.cumsum(dim=0, dtype=torch.int32) - b_q_seq_len total_token_num = prefill_input_len * batch_size b_mtp_index = torch.zeros(batch_size, dtype=torch.int32, device="cuda") + b_is_decode_req = torch.zeros(batch_size, dtype=torch.bool, device="cuda") model_input = ModelInput( batch_size=batch_size, total_token_num=total_token_num, @@ -1254,6 +1260,7 @@ def _init_padded_req(self): b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, + b_is_decode_req=b_is_decode_req, b_ready_cache_len=b_ready_cache_len, b_prefill_start_loc=b_prefill_start_loc, b_prefill_has_output_cpu=[ diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 04fe48aa25..93912dfe08 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -24,7 +24,7 @@ class ModelInput: # 在 prefill 阶段,用于在 enable_prefill_decode_mixed 开启下, # 用于标识请求是否为 decode 请求混合在 prefill 请求中。 # 其对应的 input_ids 需要特殊处理, 从 req_to_next_token_ids 中获取。 - + # 该字段是 prefill 必填参数;普通 prefill 和 draft prompt prefill 填全 False。 b_is_decode_req: torch.Tensor = None # Decode 逐行携带的 radix cache 共享长度。普通 attention backend 不使用该 @@ -36,9 +36,9 @@ class ModelInput: mem_indexes: torch.Tensor = None is_prefill: bool = False b_ready_cache_len: torch.Tensor = None - # Request/row-aligned MRoPE position offset. Normal prompt prefill leaves - # it unset; decode and one-token-per-row MTP draft KV commits carry it. - # Row-aligned input transforms must preserve the tensor unchanged. + # Request/row-aligned MRoPE position offset. It is decode-only; prefill + # builds positions directly from the complete prompt layout. + # Row-aligned decode input transforms must preserve the tensor unchanged. b_position_delta: torch.Tensor = None b_prefill_start_loc: torch.Tensor = None multimodal_params: list = None @@ -57,49 +57,76 @@ class ModelInput: def to_cuda(self): self.check_input() - if self.input_ids is not None: - self.input_ids = self.input_ids.cuda(non_blocking=True) + + # Prefill 和 decode 都必须提供的公共张量。 if self.mem_indexes is None: self.mem_indexes = self.mem_indexes_cpu.cuda(non_blocking=True) - - if self.b_is_decode_req is not None: - self.b_is_decode_req = self.b_is_decode_req.cuda(non_blocking=True) - assert self.is_prefill - self.b_req_idx = self.b_req_idx.cuda(non_blocking=True) self.b_seq_len = self.b_seq_len.cuda(non_blocking=True) self.b_mtp_index = self.b_mtp_index.cuda(non_blocking=True) - if self.b_ready_cache_len is not None: - self.b_ready_cache_len = self.b_ready_cache_len.cuda(non_blocking=True) - if self.b_position_delta is not None: - self.b_position_delta = self.b_position_delta.cuda(non_blocking=True) - else: - # Decode always needs the request-level MRoPE delta. Prefill may - # omit it for a normal prompt, while MTP draft KV commit prefill - # deliberately carries it because its rows use decode positions. - assert self.is_prefill is True, "decode ModelInput should provide b_position_delta." - if self.b_prefill_start_loc is not None: + if self.is_prefill: + # Prefill 必须提供的张量。 + self.input_ids = self.input_ids.cuda(non_blocking=True) + self.b_ready_cache_len = self.b_ready_cache_len.cuda(non_blocking=True) self.b_prefill_start_loc = self.b_prefill_start_loc.cuda(non_blocking=True) - if not self.is_prefill: + self.b_is_decode_req = self.b_is_decode_req.cuda(non_blocking=True) + else: + # Decode 必须提供的张量。 + self.b_position_delta = self.b_position_delta.cuda(non_blocking=True) self.b_shared_seq_len = self.b_shared_seq_len.cuda(non_blocking=True) self.b_shared_radix_node_id = self.b_shared_radix_node_id.cuda(non_blocking=True) + # Decode 可以显式提供 input_ids;未提供时会在模型内部按请求索引收集。 + if self.input_ids is not None: + self.input_ids = self.input_ids.cuda(non_blocking=True) + def __post_init__(self): self.check_input() def check_input(self): - assert len(self.multimodal_params) == self.batch_size if self.input_ids is not None: assert ( self.input_ids.dtype == torch.int64 ), f"model input_ids must use torch.int64, got {self.input_ids.dtype}" - if not self.is_prefill: + + assert self.b_req_idx is not None + assert self.b_mtp_index is not None + assert self.b_seq_len is not None + assert self.multimodal_params is not None + assert self.mem_indexes is not None or self.mem_indexes_cpu is not None + + assert self.b_req_idx.shape == (self.batch_size,) + assert self.b_mtp_index.shape == self.b_req_idx.shape + assert self.b_seq_len.shape == self.b_req_idx.shape + assert len(self.multimodal_params) == self.batch_size + + if self.is_prefill: + assert self.input_ids is not None + assert self.max_cache_len is not None + assert self.prefix_total_token_num is not None + assert self.b_ready_cache_len is not None + assert self.b_prefill_start_loc is not None + assert self.b_is_decode_req is not None + assert self.b_ready_cache_len.shape == self.b_req_idx.shape + assert self.b_prefill_start_loc.shape == self.b_req_idx.shape + assert self.b_is_decode_req.shape == self.b_req_idx.shape + assert self.b_is_decode_req.dtype == torch.bool + assert self.b_position_delta is None, "prefill must not provide b_position_delta" + if self.b_prefill_has_output_cpu is not None: + assert len(self.b_prefill_has_output_cpu) == self.batch_size + else: + assert self.max_q_seq_len == 1 + assert self.b_position_delta is not None assert self.b_shared_seq_len is not None assert self.b_shared_radix_node_id is not None + assert self.b_position_delta.shape == self.b_req_idx.shape assert self.b_shared_seq_len.shape == self.b_req_idx.shape assert self.b_shared_radix_node_id.shape == self.b_req_idx.shape + mem_indexes = self.mem_indexes if self.mem_indexes is not None else self.mem_indexes_cpu + assert mem_indexes.ndim == 1 + @dataclass class ModelMtpOutputCollector: @@ -136,9 +163,7 @@ def to_no_ref_tensor(self) -> None: if self.confidence_logits is not None: self.confidence_logits = tensor_to_no_ref_tensor(self.confidence_logits) - def unpad_decode( - self, padded_batch_size: int, origin_batch_size: int - ) -> "ModelMtpOutputCollector": + def unpad_decode(self, padded_batch_size: int, origin_batch_size: int) -> "ModelMtpOutputCollector": collector = copy.copy(self) if collector.spec_hidden is not None: collector.spec_hidden = collector.spec_hidden[:origin_batch_size] @@ -146,15 +171,10 @@ def unpad_decode( collector.draft_token_ids = collector.draft_token_ids[:origin_batch_size] if collector.confidence_logits is not None: confidence_row_count = collector.confidence_logits.shape[0] - assert ( - confidence_row_count > 0 - and padded_batch_size % confidence_row_count == 0 - ) + assert confidence_row_count > 0 and padded_batch_size % confidence_row_count == 0 rows_per_confidence = padded_batch_size // confidence_row_count assert origin_batch_size % rows_per_confidence == 0 - collector.confidence_logits = collector.confidence_logits[ - : origin_batch_size // rows_per_confidence - ] + collector.confidence_logits = collector.confidence_logits[: origin_batch_size // rows_per_confidence] return collector def unpad_prefill(self, origin_handle_token_num: int) -> "ModelMtpOutputCollector": diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index 5206ae5cdb..1350e358b3 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -199,6 +199,7 @@ def warmup(self, model): b_seq_len = torch.empty(1, dtype=torch.int32, device="cuda") b_seq_len.fill_(total_token_num) b_mtp_index = torch.zeros(1, dtype=torch.int32, device="cuda") + b_is_decode_req = torch.zeros(1, dtype=torch.bool, device="cuda") b_ready_cache_len = torch.zeros(1, dtype=torch.int32, device="cuda") b_prefill_start_loc = torch.zeros(1, dtype=torch.int32, device="cuda") @@ -213,6 +214,7 @@ def warmup(self, model): b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, + b_is_decode_req=b_is_decode_req, b_ready_cache_len=b_ready_cache_len, b_prefill_start_loc=b_prefill_start_loc, is_prefill=True, @@ -259,6 +261,7 @@ def warmup_overlap(self, model): b_seq_len = torch.empty(1, dtype=torch.int32, device="cuda") b_seq_len.fill_(total_token_num) b_mtp_index = torch.zeros(1, dtype=torch.int32, device="cuda") + b_is_decode_req = torch.zeros(1, dtype=torch.bool, device="cuda") b_ready_cache_len = torch.zeros(1, dtype=torch.int32, device="cuda") b_prefill_start_loc = torch.zeros(1, dtype=torch.int32, device="cuda") @@ -273,6 +276,7 @@ def warmup_overlap(self, model): b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, b_seq_len=b_seq_len, + b_is_decode_req=b_is_decode_req, b_ready_cache_len=b_ready_cache_len, b_prefill_start_loc=b_prefill_start_loc, is_prefill=True, diff --git a/lightllm/models/qwen2_vl/infer_struct.py b/lightllm/models/qwen2_vl/infer_struct.py index 399dfa866b..0c61655091 100644 --- a/lightllm/models/qwen2_vl/infer_struct.py +++ b/lightllm/models/qwen2_vl/infer_struct.py @@ -18,37 +18,17 @@ def init_some_extra_state(self, model): self.rope_type = rope_scaling.get("rope_type", rope_scaling.get("type", None)) InferStateInfo.init_some_extra_state(self, model) - # Case 1: normal prompt prefill. There is no request-level position - # delta yet, so build the complete 3-axis MRoPE positions from the - # prompt's image/video layout. - if self.is_prefill and self.b_position_delta is None: + # Prefill builds complete 3-axis MRoPE positions from the prompt's + # image/video layout. Request-level position deltas are decode-only. + if self.is_prefill: + assert self.b_position_delta is None, "prefill must not provide b_position_delta" self.position_ids = self.get_mrope_position(self.multimodal_params) - # Case 2: normal decode. Base position_ids contains one scalar position + # Decode base position_ids contains one scalar position # per request; add the cached multimodal delta and broadcast the result # to MRoPE's temporal/height/width axes. - elif not self.is_prefill: - assert self.b_position_delta is not None, "decode requires b_position_delta" - self._apply_mrope_position_delta() - - # Case 3: extend the draft-model KV cache after target verification. - # These rows originally came from target-model decode verification, - # potentially with multiple token rows belonging to the same request. - # The proposer reuses that row-aligned ModelInput and changes - # is_prefill to True so the draft model can replay all verified rows - # through its prefill kernel and write them into the draft KV cache in - # one forward. Therefore is_prefill describes the kernel/cache-write - # path here; it does not mean that this is the original prompt prefill. - # - # Every row is a post-prompt token and already carries the request-level - # b_position_delta derived from the original multimodal prompt. Its - # MRoPE position must consequently be calculated in the same way as a - # decode token: base position plus b_position_delta. Rebuilding MRoPE - # positions from multimodal_params would be incorrect because the - # original image/video token layout is no longer being prefetched and - # multimodal_params may contain only empty row-aligned placeholders. else: - assert self.b_position_delta is not None + assert self.b_position_delta is not None, "decode requires b_position_delta" self._apply_mrope_position_delta() self.position_ids = self.position_ids.contiguous() diff --git a/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py b/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py index 2e392e6ffa..67cef88409 100644 --- a/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py @@ -8,8 +8,9 @@ def prepare_mtp_prefill_inputs( b_next_token_ids: torch.Tensor, mtp_draft_input_hiddens: torch.Tensor, ) -> ModelInput: - # MTP supplies explicit token ids; mixed-prefill gathering must not replace them. - model_input.b_is_decode_req = None + # MTP supplies explicit token ids; mark every row as prefill so mixed-prefill + # gathering does not replace them with request-level decode token ids. + model_input.b_is_decode_req = torch.zeros_like(model_input.b_req_idx, dtype=torch.bool) model_input.input_ids = gen_mtp_new_input_ids( input_ids=model_input.input_ids, b_next_token_ids=b_next_token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 59e1289976..5bf941b321 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -13,7 +13,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, generate_eagle_token_ids, - prepare_eagle_verify_extend_input, + prepare_eagle_verify_decode_input, ) @@ -107,7 +107,7 @@ def propose_next_dp_eagle_autoregressive_overlap( request_capacities_by_batch.append(request_capacity) real_request_counts.append(real_request_count) accepted_tail_rows_by_batch.append(accepted_tail_rows) - prepare_eagle_verify_extend_input( + prepare_eagle_verify_decode_input( model_input=model_input, input_ids=token_ids, target_hidden=model_output.mtp_collector.spec_hidden, @@ -117,7 +117,7 @@ def propose_next_dp_eagle_autoregressive_overlap( proposal_token_ids = target_next_token_ids0.new_empty((total_real_request_count, draft_step)) draft_model = proposer.backend.draft_models[0] - extend_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) + extend_outputs = draft_model.microbatch_overlap_decode(*model_inputs) draft_token_ids_by_batch = [] draft_hiddens_by_batch = [] @@ -157,10 +157,9 @@ def propose_next_dp_eagle_autoregressive_overlap( model_input.b_req_idx = draft_req_indices_by_batch[batch_index] model_input.b_mtp_index = torch.zeros_like(model_input.b_req_idx) model_input.b_seq_len = draft_seq_lens_by_batch[batch_index] - model_input.b_position_delta = ( - position_deltas_by_batch[batch_index].index_select(0, accepted_tail_rows_by_batch[batch_index]) - if position_deltas_by_batch[batch_index] is not None - else None + assert position_deltas_by_batch[batch_index] is not None + model_input.b_position_delta = position_deltas_by_batch[batch_index].index_select( + 0, accepted_tail_rows_by_batch[batch_index] ) model_input.b_shared_seq_len = draft_shared_seq_lens_by_batch[batch_index] model_input.b_shared_radix_node_id = draft_shared_radix_node_ids_by_batch[batch_index] diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 4997236faa..945c2334af 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -59,23 +59,16 @@ def generate_eagle_token_ids_and_prob( return map_draft_token_ids(draft_token_ids), draft_token_probs -def prepare_eagle_verify_extend_input( +def prepare_eagle_verify_decode_input( model_input: ModelInput, input_ids: torch.Tensor, target_hidden: torch.Tensor, ) -> None: - model_input.is_prefill = True - model_input.total_token_num = model_input.batch_size - model_input.prefix_total_token_num = 0 - model_input.max_cache_len = max(0, int(model_input.max_cache_len or 0)) + """复用 target MTP decode 布局,为 drafter 的 verification KV commit 准备输入。""" + + assert not model_input.is_prefill model_input.input_ids = input_ids model_input.mtp_draft_input_hiddens = target_hidden - model_input.b_ready_cache_len = model_input.b_seq_len - 1 - model_input.b_prefill_start_loc = torch.arange( - model_input.batch_size, - dtype=torch.int32, - device=input_ids.device, - ) def propose_next_eagle( @@ -112,7 +105,8 @@ def propose_next_eagle( accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() draft_model = proposer.backend.draft_models[0] position_delta = target_model_input.b_position_delta - prepare_eagle_verify_extend_input( + assert position_delta is not None + prepare_eagle_verify_decode_input( model_input=target_model_input, input_ids=target_next_token_ids, target_hidden=target_model_output.mtp_collector.spec_hidden, @@ -153,13 +147,9 @@ def propose_next_eagle( draft_input.b_req_idx = target_model_input.b_req_idx.index_select(0, accepted_tail_rows) draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) draft_input.b_seq_len = draft_seq_lens - draft_input.b_position_delta = ( - position_delta.index_select(0, accepted_tail_rows) if position_delta is not None else None - ) + draft_input.b_position_delta = position_delta.index_select(0, accepted_tail_rows) draft_input.b_shared_seq_len = target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) - draft_input.b_shared_radix_node_id = target_model_input.b_shared_radix_node_id.index_select( - 0, accepted_tail_rows - ) + draft_input.b_shared_radix_node_id = target_model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) if len(draft_input.multimodal_params) != request_count: empty_multimodal_params = {"images": [], "audios": []} draft_input.multimodal_params = [empty_multimodal_params] * request_count diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py index c8824a8e23..989238ca2a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py @@ -33,17 +33,10 @@ def extend_parallel_block_draft_kv_cache( target_model_input: ModelInput, target_hidden: torch.Tensor, ) -> None: - """提交本轮 target verify hidden,扩展 parallel-block drafter KV。""" - - target_model_input.total_token_num = target_model_input.batch_size - target_model_input.prefix_total_token_num = 0 - target_model_input.is_prefill = True - target_model_input.b_ready_cache_len = target_model_input.b_seq_len - 1 - target_model_input.b_prefill_start_loc = torch.arange( - target_model_input.batch_size, - dtype=torch.int32, - device=target_hidden.device, - ) + """用 MTP decode 提交本轮 target verify hidden,扩展 drafter KV。""" + + assert not target_model_input.is_prefill + assert target_model_input.b_position_delta is not None target_model_input.mtp_draft_input_hiddens = target_hidden proposer.backend.draft_models[0].forward(target_model_input) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index cbe3d2a6ac..754dd0171a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -1,11 +1,10 @@ +import copy + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( - VanillaSpecProposal, - propose_next_chained_mtp, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal class VanillaNoAttProposer(BaseSpecProposer): @@ -28,12 +27,64 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=draft_step, - accept_len=accept_len, + req_num = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids_by_step = [] + schedule_scores_by_step = [] + + if draft_step == 0: + return VanillaSpecProposal( + token_ids=target_next_token_ids.new_empty((req_num, 0)), + extra_mem_indexes_cpu=[], + schedule_scores=( + torch.empty((req_num, 0), dtype=torch.float32, device=target_next_token_ids.device) + if self.enable_dynmaic_mtp + else None + ), + ) + + assert accept_len is not None + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + + # Vanilla No-Att 不维护 KV cache。每一级 draft 只需要每个请求本轮 + # 最后接受位置的 token 和 hidden,因此先把 target verify 布局从 + # [verify_batch_size, ...] 压缩为 [req_num, ...],后续所有 draft + # model 都只对这 req_num 行进行推理。 + draft_token_ids = target_next_token_ids.index_select(0, accepted_tail_rows) + draft_hidden = target_model_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows) + draft_input = copy.copy(target_model_input) + draft_input.batch_size = req_num + draft_input.input_ids = draft_token_ids + draft_input.b_req_idx = target_model_input.b_req_idx.index_select(0, accepted_tail_rows) + draft_input.b_mtp_index = target_model_input.b_mtp_index.index_select(0, accepted_tail_rows) + draft_input.b_seq_len = target_model_input.b_seq_len.index_select(0, accepted_tail_rows) + draft_input.mem_indexes = target_model_input.mem_indexes.index_select(0, accepted_tail_rows) + draft_input.b_shared_seq_len = target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) + draft_input.b_shared_radix_node_id = target_model_input.b_shared_radix_node_id.index_select( + 0, accepted_tail_rows + ) + draft_input.b_position_delta = target_model_input.b_position_delta.index_select(0, accepted_tail_rows) + draft_input.mem_indexes_cpu = None + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(req_num)] + + for step in range(draft_step): + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_output = self.backend.draft_models[step].forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + + if self.enable_dynmaic_mtp: + # Vanilla No-Att 没有独立的 confidence head;动态调度使用 + # 当前 draft model 选中 token 的采样概率。 + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + schedule_scores_by_step.append(draft_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + proposal_token_ids_by_step.append(draft_token_ids.unsqueeze(1)) + + proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) + schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index 30de454283..58116c9ff0 100644 --- a/unit_tests/common/basemodel/test_model_input.py +++ b/unit_tests/common/basemodel/test_model_input.py @@ -1,6 +1,9 @@ +from types import SimpleNamespace + import pytest import torch +from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.common.basemodel.batch_objs import ModelInput @@ -14,10 +17,20 @@ def _create_model_input(*, is_prefill=False): b_req_idx=torch.arange(batch_size, dtype=torch.int32), b_mtp_index=torch.zeros(batch_size, dtype=torch.int32), b_seq_len=torch.ones(batch_size, dtype=torch.int32), + mem_indexes_cpu=torch.arange(batch_size, dtype=torch.int32), is_prefill=is_prefill, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], ) - if not is_prefill: + if is_prefill: + kwargs["max_cache_len"] = 0 + kwargs["prefix_total_token_num"] = 0 + kwargs["input_ids"] = torch.ones(batch_size, dtype=torch.int64) + kwargs["b_ready_cache_len"] = torch.zeros(batch_size, dtype=torch.int32) + kwargs["b_prefill_start_loc"] = torch.arange(batch_size, dtype=torch.int32) + kwargs["b_is_decode_req"] = torch.zeros(batch_size, dtype=torch.bool) + kwargs["b_prefill_has_output_cpu"] = [False] * batch_size + else: + kwargs["b_position_delta"] = torch.zeros(batch_size, dtype=torch.int32) kwargs["b_shared_seq_len"] = torch.tensor([4, 4], dtype=torch.int32) kwargs["b_shared_radix_node_id"] = torch.tensor([10, 10], dtype=torch.int64) return ModelInput(**kwargs) @@ -33,6 +46,8 @@ def test_decode_requires_shared_radix_metadata(): b_req_idx=torch.zeros(1, dtype=torch.int32), b_mtp_index=torch.zeros(1, dtype=torch.int32), b_seq_len=torch.ones(1, dtype=torch.int32), + b_position_delta=torch.zeros(1, dtype=torch.int32), + mem_indexes_cpu=torch.zeros(1, dtype=torch.int32), is_prefill=False, multimodal_params=[{"images": [], "audios": []}], ) @@ -50,3 +65,70 @@ def test_prefill_does_not_require_shared_radix_metadata(): assert model_input.b_shared_seq_len is None assert model_input.b_shared_radix_node_id is None + + +def test_decode_requires_position_delta(): + model_input = _create_model_input() + model_input.b_position_delta = None + + with pytest.raises(AssertionError): + model_input.check_input() + + +def test_prefill_requires_prefill_metadata(): + model_input = _create_model_input(is_prefill=True) + model_input.b_ready_cache_len = None + + with pytest.raises(AssertionError): + model_input.check_input() + + +def test_prefill_requires_decode_request_markers(): + model_input = _create_model_input(is_prefill=True) + model_input.b_is_decode_req = None + + with pytest.raises(AssertionError): + model_input.check_input() + + +def test_prefill_rejects_position_delta(): + model_input = _create_model_input(is_prefill=True) + model_input.b_position_delta = torch.zeros(model_input.batch_size, dtype=torch.int32) + + with pytest.raises(AssertionError, match="prefill must not provide b_position_delta"): + model_input.check_input() + + +def test_padded_prefill_adds_non_decode_request_marker(): + model_input = ModelInput( + batch_size=1, + total_token_num=2, + max_q_seq_len=2, + max_kv_seq_len=2, + max_cache_len=0, + prefix_total_token_num=0, + input_ids=torch.ones(2, dtype=torch.int64), + b_req_idx=torch.zeros(1, dtype=torch.int32), + b_mtp_index=torch.zeros(1, dtype=torch.int32), + b_seq_len=torch.full((1,), 2, dtype=torch.int32), + b_is_decode_req=torch.ones(1, dtype=torch.bool), + b_ready_cache_len=torch.zeros(1, dtype=torch.int32), + b_prefill_start_loc=torch.zeros(1, dtype=torch.int32), + mem_indexes=torch.arange(2, dtype=torch.int32), + is_prefill=True, + b_prefill_has_output_cpu=[False], + multimodal_params=[{"images": [], "audios": []}], + ) + model = SimpleNamespace( + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=-1), + req_manager=SimpleNamespace(HOLD_REQUEST_ID=-1), + ) + + padded_input = TpPartBaseModel._create_padded_prefill_model_input( + model, + model_input=model_input, + new_handle_token_num=4, + ) + + assert padded_input.b_is_decode_req.dtype == torch.bool + assert padded_input.b_is_decode_req.tolist() == [True, False] diff --git a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py index 42c4a44f68..2cb95f5923 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py @@ -123,6 +123,7 @@ def test_compaction_preserves_shared_radix_metadata(): b_req_idx=torch.tensor([0, 0, 0, 0, 0], dtype=torch.int32, device="cuda"), b_mtp_index=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device="cuda"), b_seq_len=torch.tensor([3, 4, 5, 6, 7], dtype=torch.int32, device="cuda"), + b_position_delta=torch.zeros(5, dtype=torch.int32, device="cuda"), b_shared_seq_len=torch.full((5,), 7, dtype=torch.int32, device="cuda"), b_shared_radix_node_id=torch.full((5,), 10, dtype=torch.int64, device="cuda"), mem_indexes=torch.arange(5, dtype=torch.int32, device="cuda"), diff --git a/unit_tests/models/qwen2_vl/test_infer_struct.py b/unit_tests/models/qwen2_vl/test_infer_struct.py index e5aa40f4a8..b239e64067 100644 --- a/unit_tests/models/qwen2_vl/test_infer_struct.py +++ b/unit_tests/models/qwen2_vl/test_infer_struct.py @@ -3,7 +3,6 @@ import pytest import torch -from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.basemodel.infer_struct import InferStateInfo from lightllm.models.qwen2_vl.infer_struct import Qwen2VLInferStateInfo @@ -57,7 +56,7 @@ def test_normal_decode_applies_position_delta(monkeypatch): assert torch.equal(infer_state.position_ids, expected_position_ids) -def test_draft_commit_prefill_uses_existing_position_delta(monkeypatch): +def test_prefill_rejects_position_delta(monkeypatch): _patch_base_position_ids(monkeypatch) infer_state = Qwen2VLInferStateInfo() @@ -65,29 +64,5 @@ def test_draft_commit_prefill_uses_existing_position_delta(monkeypatch): infer_state.b_position_delta = torch.tensor([3, 5], dtype=torch.int32) infer_state.multimodal_params = [{"images": [], "audios": []}] * 2 - infer_state.init_some_extra_state(_make_model()) - - expected_position_ids = torch.tensor([[12, 24]] * 3, dtype=torch.int32) - assert torch.equal(infer_state.position_ids, expected_position_ids) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") -def test_draft_commit_prefill_model_input_accepts_position_delta(): - model_input = ModelInput( - batch_size=2, - total_token_num=2, - max_q_seq_len=1, - max_kv_seq_len=8, - input_ids=torch.tensor([1, 2], dtype=torch.int64), - b_req_idx=torch.tensor([0, 1], dtype=torch.int32), - b_mtp_index=torch.zeros(2, dtype=torch.int32), - b_seq_len=torch.tensor([4, 6], dtype=torch.int32), - b_position_delta=torch.tensor([3, 5], dtype=torch.int32), - mem_indexes_cpu=torch.tensor([10, 11], dtype=torch.int32), - is_prefill=True, - multimodal_params=[{"images": [], "audios": []}] * 2, - ) - - model_input.to_cuda() - - assert model_input.b_position_delta.is_cuda + with pytest.raises(AssertionError, match="prefill must not provide b_position_delta"): + infer_state.init_some_extra_state(_make_model()) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index c8a59c58c1..77def5b36a 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -156,13 +156,12 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): draft_step=2, ) - assert draft_model.extend_inputs[0] is model_input0 - assert draft_model.extend_inputs[1] is model_input1 - assert len(draft_model.decode_inputs) == 1 + assert draft_model.extend_inputs is None + assert len(draft_model.decode_inputs) == 2 assert draft_model.decode_inputs[0][0] is model_input0 assert draft_model.decode_inputs[0][1] is model_input1 - assert draft_model.extend_batch_sizes == (6, 6) - assert draft_model.decode_batch_sizes == [(2, 2)] + assert draft_model.extend_batch_sizes is None + assert draft_model.decode_batch_sizes == [(6, 6), (2, 2)] assert proposal.token_ids.shape == (3, 2) assert torch.equal(proposal.token_ids, torch.tensor([[1, 0], [0, 0], [5, 1]])) assert len(proposal.extra_mem_indexes_cpu) == 1 diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py new file mode 100644 index 0000000000..35c08b2af5 --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py @@ -0,0 +1,106 @@ +from types import SimpleNamespace + +import torch + +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer + + +def test_vanilla_no_att_proposes_from_one_accepted_tail_per_request(): + draft_calls = [] + draft_outputs = [ + SimpleNamespace( + token_ids=torch.tensor([21, 24]), + token_probs=torch.tensor([0.8, 0.7]), + mtp_collector=SimpleNamespace(spec_hidden=torch.tensor([[102.0, 103.0], [108.0, 109.0]])), + ), + SimpleNamespace( + token_ids=torch.tensor([31, 34]), + token_probs=torch.tensor([0.6, 0.5]), + mtp_collector=SimpleNamespace(spec_hidden=torch.tensor([[202.0, 203.0], [208.0, 209.0]])), + ), + ] + + def build_draft_model(step): + def forward(model_input): + draft_calls.append( + { + "batch_size": model_input.batch_size, + "input_ids": model_input.input_ids.clone(), + "draft_hidden": model_input.mtp_draft_input_hiddens.clone(), + "b_req_idx": model_input.b_req_idx.clone(), + "b_mtp_index": model_input.b_mtp_index.clone(), + "b_seq_len": model_input.b_seq_len.clone(), + "mem_indexes": model_input.mem_indexes.clone(), + } + ) + return draft_outputs[step] + + return SimpleNamespace(forward=forward) + + backend = SimpleNamespace( + draft_models=[build_draft_model(0), build_draft_model(1)], + _gen_argmax_token_ids=lambda output: output.token_ids, + _gen_argmax_token_ids_and_prob=lambda output: (output.token_ids, output.token_probs), + ) + proposer = VanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=True) + target_model_input = SimpleNamespace( + batch_size=5, + input_ids=torch.tensor([10, 11, 12, 13, 14]), + b_req_idx=torch.tensor([7, 7, 7, 9, 9], dtype=torch.int32), + b_mtp_index=torch.tensor([0, 1, 2, 0, 1], dtype=torch.int32), + b_seq_len=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32), + mem_indexes=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32), + mem_indexes_cpu=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32), + b_position_delta=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32), + b_shared_seq_len=torch.tensor([8, 8, 8, 6, 6], dtype=torch.int32), + b_shared_radix_node_id=torch.tensor([70, 70, 70, 90, 90], dtype=torch.int64), + multimodal_params=[{"images": [], "audios": []} for _ in range(5)], + ) + target_model_output = SimpleNamespace( + mtp_collector=SimpleNamespace(spec_hidden=torch.arange(10, dtype=torch.float32).reshape(5, 2)) + ) + + proposal = proposer.propose_next( + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_model_input.input_ids, + b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32), + draft_step=2, + accept_len=torch.tensor([2, 2], dtype=torch.int32), + ) + + torch.testing.assert_close(proposal.token_ids, torch.tensor([[21, 31], [24, 34]])) + torch.testing.assert_close(proposal.schedule_scores, torch.tensor([[0.8, 0.6], [0.7, 0.5]])) + assert len(draft_calls) == 2 + assert draft_calls[0]["batch_size"] == 2 + torch.testing.assert_close(draft_calls[0]["input_ids"], torch.tensor([11, 14])) + torch.testing.assert_close(draft_calls[0]["draft_hidden"], torch.tensor([[2.0, 3.0], [8.0, 9.0]])) + torch.testing.assert_close(draft_calls[0]["b_req_idx"], torch.tensor([7, 9], dtype=torch.int32)) + torch.testing.assert_close(draft_calls[0]["b_mtp_index"], torch.tensor([1, 1], dtype=torch.int32)) + torch.testing.assert_close(draft_calls[0]["b_seq_len"], torch.tensor([11, 21], dtype=torch.int32)) + torch.testing.assert_close(draft_calls[0]["mem_indexes"], torch.tensor([101, 104], dtype=torch.int32)) + torch.testing.assert_close(draft_calls[1]["input_ids"], torch.tensor([21, 24])) + torch.testing.assert_close( + draft_calls[1]["draft_hidden"], + torch.tensor([[102.0, 103.0], [108.0, 109.0]]), + ) + assert target_model_input.batch_size == 5 + torch.testing.assert_close(target_model_input.input_ids, torch.tensor([10, 11, 12, 13, 14])) + + +def test_vanilla_no_att_skips_draft_forward_for_zero_steps(): + proposer = VanillaNoAttProposer( + backend=SimpleNamespace(draft_models=[]), + enable_dynmaic_mtp=True, + ) + + proposal = proposer.propose_next( + target_model_input=None, + target_model_output=None, + target_next_token_ids=torch.tensor([10, 11]), + b_req_mtp_start_loc=torch.tensor([0, 1], dtype=torch.int32), + draft_step=0, + ) + + assert proposal.token_ids.shape == (2, 0) + assert proposal.schedule_scores.shape == (2, 0) diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index 711e9505a8..ba2afdfd09 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -17,6 +17,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer, DFlashSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( build_parallel_block_draft_input, + extend_parallel_block_draft_kv_cache, ) from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.utils import envs_utils @@ -332,17 +333,19 @@ def test_dflash_reuses_decode_input_for_kv_commit(): ) target_hidden = torch.empty(4, 16) - proposer.extend_draft_kv_cache(model_input, target_hidden) + extend_parallel_block_draft_kv_cache( + proposer=proposer, + target_model_input=model_input, + target_hidden=target_hidden, + ) assert len(forwarded_inputs) == 1 assert forwarded_inputs[0] is model_input assert model_input.input_ids is input_ids assert model_input.multimodal_params is multimodal_params - assert model_input.total_token_num == model_input.batch_size - assert model_input.prefix_total_token_num == 0 - assert model_input.is_prefill - torch.testing.assert_close(model_input.b_ready_cache_len, model_input.b_seq_len - 1) - torch.testing.assert_close(model_input.b_prefill_start_loc, torch.arange(4, dtype=torch.int32)) + assert model_input.total_token_num == 20 + assert model_input.prefix_total_token_num is None + assert not model_input.is_prefill assert model_input.b_position_delta is position_delta assert model_input.mtp_draft_input_hiddens is target_hidden From 8aead20edca1e7ba51c973279944871cc0669b3a Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 07:03:39 +0000 Subject: [PATCH 066/103] perf: fuse vanilla MTP row selection --- .../triton_kernel/select_mtp_rows.py | 185 ++++++++++++++++++ .../proposers/vanilla_no_att.py | 36 ++-- .../mtp_speculative/test_vanilla_no_att.py | 87 ++++++-- 3 files changed, 277 insertions(+), 31 deletions(-) create mode 100644 lightllm/common/basemodel/triton_kernel/select_mtp_rows.py diff --git a/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py b/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py new file mode 100644 index 0000000000..1ed2877f12 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py @@ -0,0 +1,185 @@ +"""MTP verification row-selection kernels.""" + +from typing import NamedTuple + +import torch +import triton +import triton.language as tl + + +class SelectedMtpRows(NamedTuple): + input_ids: torch.Tensor + hidden: torch.Tensor + b_req_idx: torch.Tensor + b_mtp_index: torch.Tensor + b_seq_len: torch.Tensor + mem_indexes: torch.Tensor + b_shared_seq_len: torch.Tensor + b_shared_radix_node_id: torch.Tensor + b_position_delta: torch.Tensor + + +@triton.jit +def _select_accepted_tail_rows_kernel( + b_req_mtp_start_loc, + accept_len, + input_ids, + input_ids_stride, + hidden, + hidden_stride_0, + hidden_stride_1, + b_req_idx, + b_req_idx_stride, + b_mtp_index, + b_mtp_index_stride, + b_seq_len, + b_seq_len_stride, + mem_indexes, + mem_indexes_stride, + b_shared_seq_len, + b_shared_seq_len_stride, + b_shared_radix_node_id, + b_shared_radix_node_id_stride, + b_position_delta, + b_position_delta_stride, + out_input_ids, + out_hidden, + out_hidden_stride_0, + out_hidden_stride_1, + out_b_req_idx, + out_b_mtp_index, + out_b_seq_len, + out_mem_indexes, + out_b_shared_seq_len, + out_b_shared_radix_node_id, + out_b_position_delta, + hidden_size, + BLOCK_HIDDEN: tl.constexpr, + PIPELINE_STAGES: tl.constexpr, +): + out_row = tl.program_id(0) + src_row = tl.load(b_req_mtp_start_loc + out_row) + tl.load(accept_len + out_row) - 1 + + tl.store(out_input_ids + out_row, tl.load(input_ids + src_row * input_ids_stride)) + tl.store(out_b_req_idx + out_row, tl.load(b_req_idx + src_row * b_req_idx_stride)) + tl.store(out_b_mtp_index + out_row, tl.load(b_mtp_index + src_row * b_mtp_index_stride)) + tl.store(out_b_seq_len + out_row, tl.load(b_seq_len + src_row * b_seq_len_stride)) + tl.store(out_mem_indexes + out_row, tl.load(mem_indexes + src_row * mem_indexes_stride)) + tl.store( + out_b_shared_seq_len + out_row, + tl.load(b_shared_seq_len + src_row * b_shared_seq_len_stride), + ) + tl.store( + out_b_shared_radix_node_id + out_row, + tl.load(b_shared_radix_node_id + src_row * b_shared_radix_node_id_stride), + ) + tl.store( + out_b_position_delta + out_row, + tl.load(b_position_delta + src_row * b_position_delta_stride), + ) + + hidden_block_offsets = tl.arange(0, BLOCK_HIDDEN) + for hidden_start in tl.range(0, hidden_size, BLOCK_HIDDEN, num_stages=PIPELINE_STAGES): + hidden_offsets = hidden_start + hidden_block_offsets + hidden_mask = hidden_offsets < hidden_size + hidden_values = tl.load( + hidden + src_row * hidden_stride_0 + hidden_offsets * hidden_stride_1, + mask=hidden_mask, + other=0, + ) + tl.store( + out_hidden + out_row * out_hidden_stride_0 + hidden_offsets * out_hidden_stride_1, + hidden_values, + mask=hidden_mask, + ) + + +@torch.no_grad() +def select_accepted_tail_rows( + b_req_mtp_start_loc: torch.Tensor, + accept_len: torch.Tensor, + input_ids: torch.Tensor, + hidden: torch.Tensor, + b_req_idx: torch.Tensor, + b_mtp_index: torch.Tensor, + b_seq_len: torch.Tensor, + mem_indexes: torch.Tensor, + b_shared_seq_len: torch.Tensor, + b_shared_radix_node_id: torch.Tensor, + b_position_delta: torch.Tensor, +) -> SelectedMtpRows: + """Select one accepted-tail row per request in a single CUDA kernel.""" + + req_num = b_req_mtp_start_loc.shape[0] + assert req_num > 0 + assert input_ids.is_cuda + assert hidden.ndim == 2 and hidden.shape[0] == input_ids.shape[0] + assert hidden.shape[1] > 0 + tensors = ( + b_req_mtp_start_loc, + accept_len, + hidden, + b_req_idx, + b_mtp_index, + b_seq_len, + mem_indexes, + b_shared_seq_len, + b_shared_radix_node_id, + b_position_delta, + ) + assert all(tensor.is_cuda and tensor.device == input_ids.device for tensor in tensors) + + selected = SelectedMtpRows( + input_ids=input_ids.new_empty((req_num,)), + hidden=hidden.new_empty((req_num, hidden.shape[1])), + b_req_idx=b_req_idx.new_empty((req_num,)), + b_mtp_index=b_mtp_index.new_empty((req_num,)), + b_seq_len=b_seq_len.new_empty((req_num,)), + mem_indexes=mem_indexes.new_empty((req_num,)), + b_shared_seq_len=b_shared_seq_len.new_empty((req_num,)), + b_shared_radix_node_id=b_shared_radix_node_id.new_empty((req_num,)), + b_position_delta=b_position_delta.new_empty((req_num,)), + ) + block_hidden = 1024 + pipeline_stages = 3 + grid = (req_num,) + _select_accepted_tail_rows_kernel[grid]( + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + input_ids=input_ids, + input_ids_stride=input_ids.stride(0), + hidden=hidden, + hidden_stride_0=hidden.stride(0), + hidden_stride_1=hidden.stride(1), + b_req_idx=b_req_idx, + b_req_idx_stride=b_req_idx.stride(0), + b_mtp_index=b_mtp_index, + b_mtp_index_stride=b_mtp_index.stride(0), + b_seq_len=b_seq_len, + b_seq_len_stride=b_seq_len.stride(0), + mem_indexes=mem_indexes, + mem_indexes_stride=mem_indexes.stride(0), + b_shared_seq_len=b_shared_seq_len, + b_shared_seq_len_stride=b_shared_seq_len.stride(0), + b_shared_radix_node_id=b_shared_radix_node_id, + b_shared_radix_node_id_stride=b_shared_radix_node_id.stride(0), + b_position_delta=b_position_delta, + b_position_delta_stride=b_position_delta.stride(0), + out_input_ids=selected.input_ids, + out_hidden=selected.hidden, + out_hidden_stride_0=selected.hidden.stride(0), + out_hidden_stride_1=selected.hidden.stride(1), + out_b_req_idx=selected.b_req_idx, + out_b_mtp_index=selected.b_mtp_index, + out_b_seq_len=selected.b_seq_len, + out_mem_indexes=selected.mem_indexes, + out_b_shared_seq_len=selected.b_shared_seq_len, + out_b_shared_radix_node_id=selected.b_shared_radix_node_id, + out_b_position_delta=selected.b_position_delta, + hidden_size=hidden.shape[1], + BLOCK_HIDDEN=block_hidden, + PIPELINE_STAGES=pipeline_stages, + num_warps=8, + num_stages=pipeline_stages, + ) + return selected diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index 754dd0171a..396720b834 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -3,6 +3,7 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal @@ -43,26 +44,35 @@ def propose_next( ) assert accept_len is not None - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - # Vanilla No-Att 不维护 KV cache。每一级 draft 只需要每个请求本轮 # 最后接受位置的 token 和 hidden,因此先把 target verify 布局从 # [verify_batch_size, ...] 压缩为 [req_num, ...],后续所有 draft # model 都只对这 req_num 行进行推理。 - draft_token_ids = target_next_token_ids.index_select(0, accepted_tail_rows) - draft_hidden = target_model_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows) + selected_rows = select_accepted_tail_rows( + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + input_ids=target_next_token_ids, + hidden=target_model_output.mtp_collector.spec_hidden, + b_req_idx=target_model_input.b_req_idx, + b_mtp_index=target_model_input.b_mtp_index, + b_seq_len=target_model_input.b_seq_len, + mem_indexes=target_model_input.mem_indexes, + b_shared_seq_len=target_model_input.b_shared_seq_len, + b_shared_radix_node_id=target_model_input.b_shared_radix_node_id, + b_position_delta=target_model_input.b_position_delta, + ) + draft_token_ids = selected_rows.input_ids + draft_hidden = selected_rows.hidden draft_input = copy.copy(target_model_input) draft_input.batch_size = req_num draft_input.input_ids = draft_token_ids - draft_input.b_req_idx = target_model_input.b_req_idx.index_select(0, accepted_tail_rows) - draft_input.b_mtp_index = target_model_input.b_mtp_index.index_select(0, accepted_tail_rows) - draft_input.b_seq_len = target_model_input.b_seq_len.index_select(0, accepted_tail_rows) - draft_input.mem_indexes = target_model_input.mem_indexes.index_select(0, accepted_tail_rows) - draft_input.b_shared_seq_len = target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) - draft_input.b_shared_radix_node_id = target_model_input.b_shared_radix_node_id.index_select( - 0, accepted_tail_rows - ) - draft_input.b_position_delta = target_model_input.b_position_delta.index_select(0, accepted_tail_rows) + draft_input.b_req_idx = selected_rows.b_req_idx + draft_input.b_mtp_index = selected_rows.b_mtp_index + draft_input.b_seq_len = selected_rows.b_seq_len + draft_input.mem_indexes = selected_rows.mem_indexes + draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len + draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id + draft_input.b_position_delta = selected_rows.b_position_delta draft_input.mem_indexes_cpu = None draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(req_num)] diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py index 35c08b2af5..d2ffb793e3 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py @@ -1,11 +1,62 @@ from types import SimpleNamespace +import pytest import torch +from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_select_accepted_tail_rows_triton_matches_index_select(): + device = "cuda" + b_req_mtp_start_loc = torch.tensor([0, 3, 6], dtype=torch.int32, device=device) + accept_len = torch.tensor([2, 3, 2], dtype=torch.int32, device=device) + expected_rows = torch.tensor([1, 5, 7], dtype=torch.int64, device=device) + + # Exercise non-contiguous source strides as well as hidden widths spanning + # multiple Triton column blocks. + input_ids = torch.arange(16, dtype=torch.int64, device=device)[::2] + hidden = torch.arange(8 * 514, dtype=torch.float32, device=device).reshape(8, 514)[:, ::2] + b_req_idx = (torch.arange(16, dtype=torch.int32, device=device) + 10)[::2] + b_mtp_index = torch.arange(8, dtype=torch.int32, device=device) + b_seq_len = (torch.arange(16, dtype=torch.int32, device=device) + 20)[::2] + mem_indexes = (torch.arange(16, dtype=torch.int32, device=device) + 100)[::2] + b_shared_seq_len = (torch.arange(16, dtype=torch.int32, device=device) + 30)[::2] + b_shared_radix_node_id = (torch.arange(16, dtype=torch.int64, device=device) + 1000)[::2] + b_position_delta = (torch.arange(16, dtype=torch.int32, device=device) - 8)[::2] + + selected = select_accepted_tail_rows( + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + input_ids=input_ids, + hidden=hidden, + b_req_idx=b_req_idx, + b_mtp_index=b_mtp_index, + b_seq_len=b_seq_len, + mem_indexes=mem_indexes, + b_shared_seq_len=b_shared_seq_len, + b_shared_radix_node_id=b_shared_radix_node_id, + b_position_delta=b_position_delta, + ) + + torch.testing.assert_close(selected.input_ids, input_ids.index_select(0, expected_rows)) + torch.testing.assert_close(selected.hidden, hidden.index_select(0, expected_rows)) + torch.testing.assert_close(selected.b_req_idx, b_req_idx.index_select(0, expected_rows)) + torch.testing.assert_close(selected.b_mtp_index, b_mtp_index.index_select(0, expected_rows)) + torch.testing.assert_close(selected.b_seq_len, b_seq_len.index_select(0, expected_rows)) + torch.testing.assert_close(selected.mem_indexes, mem_indexes.index_select(0, expected_rows)) + torch.testing.assert_close(selected.b_shared_seq_len, b_shared_seq_len.index_select(0, expected_rows)) + torch.testing.assert_close( + selected.b_shared_radix_node_id, + b_shared_radix_node_id.index_select(0, expected_rows), + ) + torch.testing.assert_close(selected.b_position_delta, b_position_delta.index_select(0, expected_rows)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_vanilla_no_att_proposes_from_one_accepted_tail_per_request(): + device = "cuda" draft_calls = [] draft_outputs = [ SimpleNamespace( @@ -25,12 +76,12 @@ def forward(model_input): draft_calls.append( { "batch_size": model_input.batch_size, - "input_ids": model_input.input_ids.clone(), - "draft_hidden": model_input.mtp_draft_input_hiddens.clone(), - "b_req_idx": model_input.b_req_idx.clone(), - "b_mtp_index": model_input.b_mtp_index.clone(), - "b_seq_len": model_input.b_seq_len.clone(), - "mem_indexes": model_input.mem_indexes.clone(), + "input_ids": model_input.input_ids.cpu(), + "draft_hidden": model_input.mtp_draft_input_hiddens.cpu(), + "b_req_idx": model_input.b_req_idx.cpu(), + "b_mtp_index": model_input.b_mtp_index.cpu(), + "b_seq_len": model_input.b_seq_len.cpu(), + "mem_indexes": model_input.mem_indexes.cpu(), } ) return draft_outputs[step] @@ -45,28 +96,28 @@ def forward(model_input): proposer = VanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=True) target_model_input = SimpleNamespace( batch_size=5, - input_ids=torch.tensor([10, 11, 12, 13, 14]), - b_req_idx=torch.tensor([7, 7, 7, 9, 9], dtype=torch.int32), - b_mtp_index=torch.tensor([0, 1, 2, 0, 1], dtype=torch.int32), - b_seq_len=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32), - mem_indexes=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32), + input_ids=torch.tensor([10, 11, 12, 13, 14], device=device), + b_req_idx=torch.tensor([7, 7, 7, 9, 9], dtype=torch.int32, device=device), + b_mtp_index=torch.tensor([0, 1, 2, 0, 1], dtype=torch.int32, device=device), + b_seq_len=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32, device=device), + mem_indexes=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32, device=device), mem_indexes_cpu=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32), - b_position_delta=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32), - b_shared_seq_len=torch.tensor([8, 8, 8, 6, 6], dtype=torch.int32), - b_shared_radix_node_id=torch.tensor([70, 70, 70, 90, 90], dtype=torch.int64), + b_position_delta=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device=device), + b_shared_seq_len=torch.tensor([8, 8, 8, 6, 6], dtype=torch.int32, device=device), + b_shared_radix_node_id=torch.tensor([70, 70, 70, 90, 90], dtype=torch.int64, device=device), multimodal_params=[{"images": [], "audios": []} for _ in range(5)], ) target_model_output = SimpleNamespace( - mtp_collector=SimpleNamespace(spec_hidden=torch.arange(10, dtype=torch.float32).reshape(5, 2)) + mtp_collector=SimpleNamespace(spec_hidden=torch.arange(10, dtype=torch.float32, device=device).reshape(5, 2)) ) proposal = proposer.propose_next( target_model_input=target_model_input, target_model_output=target_model_output, target_next_token_ids=target_model_input.input_ids, - b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32), + b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32, device=device), draft_step=2, - accept_len=torch.tensor([2, 2], dtype=torch.int32), + accept_len=torch.tensor([2, 2], dtype=torch.int32, device=device), ) torch.testing.assert_close(proposal.token_ids, torch.tensor([[21, 31], [24, 34]])) @@ -85,7 +136,7 @@ def forward(model_input): torch.tensor([[102.0, 103.0], [108.0, 109.0]]), ) assert target_model_input.batch_size == 5 - torch.testing.assert_close(target_model_input.input_ids, torch.tensor([10, 11, 12, 13, 14])) + torch.testing.assert_close(target_model_input.input_ids.cpu(), torch.tensor([10, 11, 12, 13, 14])) def test_vanilla_no_att_skips_draft_forward_for_zero_steps(): From 69768e01f0d8bdc0301355df8e39e17391f5fb8d Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 07:34:41 +0000 Subject: [PATCH 067/103] refactor: localize MTP prefill state preparation --- .../mode_backend/mtp_pre_process.py | 21 --- .../dp_overlap_proposers/eagle_utils.py | 7 +- .../dp_overlap_proposers/vanilla_no_att.py | 14 +- .../dp_overlap_proposers/vanilla_utils.py | 17 ++- .../dp_overlap_proposers/vanilla_with_att.py | 86 ++++++++++- .../dp_proposers/vanilla_no_att.py | 3 +- .../dp_proposers/vanilla_with_att.py | 86 ++++++++++- .../mtp_speculative/proposers/eagle_utils.py | 27 +++- .../proposers/vanilla_utils.py | 23 --- .../proposers/vanilla_with_att.py | 93 +++++++++++- .../mtp_speculative/test_vanilla_no_att.py | 18 +++ .../mtp_speculative/test_vanilla_prefill.py | 142 ++++++++++++++++++ 12 files changed, 457 insertions(+), 80 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py create mode 100644 unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py diff --git a/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py b/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py deleted file mode 100644 index 67cef88409..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py +++ /dev/null @@ -1,21 +0,0 @@ -import torch -from lightllm.common.basemodel.batch_objs import ModelInput -from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids - - -def prepare_mtp_prefill_inputs( - model_input: ModelInput, - b_next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: torch.Tensor, -) -> ModelInput: - # MTP supplies explicit token ids; mark every row as prefill so mixed-prefill - # gathering does not replace them with request-level decode token ids. - model_input.b_is_decode_req = torch.zeros_like(model_input.b_req_idx, dtype=torch.bool) - model_input.input_ids = gen_mtp_new_input_ids( - input_ids=model_input.input_ids, - b_next_token_ids=b_next_token_ids, - b_seq_len=model_input.b_seq_len, - b_ready_cache_len=model_input.b_ready_cache_len, - ) - model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens - return model_input diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 5bf941b321..e086b992f7 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -12,6 +12,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( EagleSpecProposal, + _prepare_eagle_prefill_inputs, generate_eagle_token_ids, prepare_eagle_verify_decode_input, ) @@ -28,14 +29,12 @@ def fill_dp_eagle_draft_model_kv_state_overlap( ) -> None: """使用两个 target prefill microbatch 初始化 EAGLE draft state。""" - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - prepare_mtp_prefill_inputs( + _prepare_eagle_prefill_inputs( model_input=target_model_input0, b_next_token_ids=target_next_token_ids0, mtp_draft_input_hiddens=target_model_output0.mtp_collector.spec_hidden, ) - prepare_mtp_prefill_inputs( + _prepare_eagle_prefill_inputs( model_input=target_model_input1, b_next_token_ids=target_next_token_ids1, mtp_draft_input_hiddens=target_model_output1.mtp_collector.spec_hidden, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index bb0384b477..d70f780f54 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -3,12 +3,10 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( - fill_dp_chained_mtp_draft_model_kv_state_overlap, propose_next_dp_chained_mtp_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -22,7 +20,7 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) + pass def fill_draft_model_kv_state_overlap( self, @@ -33,15 +31,7 @@ def fill_draft_model_kv_state_overlap( target_model_output1: ModelOutput, target_next_token_ids1: torch.Tensor, ) -> None: - fill_dp_chained_mtp_draft_model_kv_state_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - ) + pass def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py index fa26f7c8c6..328d6a52d8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py @@ -2,6 +2,8 @@ from __future__ import annotations +import copy + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput @@ -20,18 +22,25 @@ def fill_dp_chained_mtp_draft_model_kv_state_overlap( ) -> None: """为两个 DP microbatch 依次构建 Vanilla chained draft state。""" - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs + target_model_inputs = (target_model_input0, target_model_input1) + target_next_token_ids = (target_next_token_ids0, target_next_token_ids1) + for model_input, next_token_ids in zip(target_model_inputs, target_next_token_ids): + assert model_input.is_prefill + assert model_input.b_position_delta is None + assert next_token_ids.shape == model_input.b_req_idx.shape - model_inputs = (target_model_input0, target_model_input1) + # 两个 microbatch 各自保留逐级左移后的 draft 输入,不修改 + # target prefill ModelInput。 + model_inputs = [copy.copy(model_input) for model_input in target_model_inputs] draft_hiddens = [ target_model_output0.mtp_collector.spec_hidden, target_model_output1.mtp_collector.spec_hidden, ] - draft_token_ids = [target_next_token_ids0, target_next_token_ids1] + draft_token_ids = list(target_next_token_ids) for draft_model in proposer.backend.draft_models: for batch_index, model_input in enumerate(model_inputs): - prepare_mtp_prefill_inputs( + model_inputs[batch_index] = proposer._prepare_mtp_prefill_inputs( model_input=model_input, b_next_token_ids=draft_token_ids[batch_index], mtp_draft_input_hiddens=draft_hiddens[batch_index], diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 3d3d301c8d..2396db8187 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -1,6 +1,9 @@ +import copy + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( fill_dp_chained_mtp_draft_model_kv_state_overlap, @@ -8,9 +11,9 @@ ) from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager class DpOverlapVanillaWithAttProposer(BaseDpOverlapProposer): @@ -22,7 +25,22 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) + assert target_model_input.is_prefill + assert target_model_input.b_position_delta is None + assert target_next_token_ids.shape == target_model_input.b_req_idx.shape + + draft_input = copy.copy(target_model_input) + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_token_ids = target_next_token_ids + for draft_model in self.backend.draft_models: + draft_input = self._prepare_mtp_prefill_inputs( + model_input=draft_input, + b_next_token_ids=draft_token_ids, + mtp_draft_input_hiddens=draft_hidden, + ) + draft_output = draft_model.forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) def fill_draft_model_kv_state_overlap( self, @@ -90,3 +108,67 @@ def propose_next_overlap( accept_len1, draft_step, ) + + def _prepare_mtp_prefill_inputs( + self, + model_input: ModelInput, + b_next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, + ) -> ModelInput: + """构造 Vanilla chained MTP 下一层 draft model 的 prefill 输入。 + + 每个请求当前参与计算的 query 长度为 + ``b_seq_len - b_ready_cache_len``。本方法在各请求自己的 query + 区间内将 token 左移一位,即丢弃区间首 token,并在区间尾部追加该 + 请求的 ``b_next_token_ids``。这样,下一层 MTP model 看到的 token + 与上一层 model 生成的 next token 连续对齐。Chunked prefill 时只 + 移动当前 chunk,已经写入 KV cache 的 prefix 不参与移动。 + + 同时,本方法会完成两项辅助输入设置: + + 1. 将 ``b_is_decode_req`` 设置为全 False 的缓存 GPU 常量张量,确保 + mixed-prefill 逻辑保留这里显式生成的 token,而不会按 decode + 请求重新收集 token。 + 2. 将上一层 model 输出的 hidden states 绑定到 + ``mtp_draft_input_hiddens``,供下一层 MTP model 使用。 + + 参数: + model_input: 当前层的 prefill 输入。方法会原地更新该对象的 + ``input_ids``、``b_is_decode_req`` 和 + ``mtp_draft_input_hiddens`` 字段。 + b_next_token_ids: 每个请求需要追加到当前 query 尾部的 token, + shape 为 ``[batch_size]``。 + mtp_draft_input_hiddens: 上一层 model 产生、供下一层 draft model + 使用的 hidden states。 + + 返回: + 更新后的 ``model_input``,与传入对象是同一个对象。 + + 示例: + 假设 ``b_seq_len=[4, 5]``、``b_ready_cache_len=[1, 2]``,则两个 + 请求当前 query 长度均为 3。若扁平输入和追加 token 为: + + ``input_ids = [10, 11, 12, 20, 21, 22]`` + ``b_next_token_ids = [13, 23]`` + + 更新后的扁平输入为: + + ``input_ids = [11, 12, 13, 21, 22, 23]`` + + 两个请求已缓存的 prefix 长度分别为 1 和 2,它们对应的 token + 不在 ``input_ids`` 中,因此不会被本方法移动或重写。 + """ + model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( + key="dp_overlap_vanilla_mtp_prefill_b_is_decode_req", + shape=model_input.b_req_idx.shape, + fill_value=False, + dtype=torch.bool, + ) + model_input.input_ids = gen_mtp_new_input_ids( + input_ids=model_input.input_ids, + b_next_token_ids=b_next_token_ids, + b_seq_len=model_input.b_seq_len, + b_ready_cache_len=model_input.b_ready_cache_len, + ) + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + return model_input diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py index ae51ab3219..c7d4ec380c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -4,7 +4,6 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) @@ -18,7 +17,7 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) + pass def propose_next( self, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py index cb829c44fd..9c150a4580 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -1,12 +1,15 @@ +import copy + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager class DpVanillaWithAttProposer(BaseDpProposer): @@ -18,7 +21,22 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) + assert target_model_input.is_prefill + assert target_model_input.b_position_delta is None + assert target_next_token_ids.shape == target_model_input.b_req_idx.shape + + draft_input = copy.copy(target_model_input) + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_token_ids = target_next_token_ids + for draft_model in self.backend.draft_models: + draft_input = self._prepare_mtp_prefill_inputs( + model_input=draft_input, + b_next_token_ids=draft_token_ids, + mtp_draft_input_hiddens=draft_hidden, + ) + draft_output = draft_model.forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) def propose_next( self, @@ -38,3 +56,67 @@ def propose_next( draft_step, accept_len, ) + + def _prepare_mtp_prefill_inputs( + self, + model_input: ModelInput, + b_next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, + ) -> ModelInput: + """构造 Vanilla chained MTP 下一层 draft model 的 prefill 输入。 + + 每个请求当前参与计算的 query 长度为 + ``b_seq_len - b_ready_cache_len``。本方法在各请求自己的 query + 区间内将 token 左移一位,即丢弃区间首 token,并在区间尾部追加该 + 请求的 ``b_next_token_ids``。这样,下一层 MTP model 看到的 token + 与上一层 model 生成的 next token 连续对齐。Chunked prefill 时只 + 移动当前 chunk,已经写入 KV cache 的 prefix 不参与移动。 + + 同时,本方法会完成两项辅助输入设置: + + 1. 将 ``b_is_decode_req`` 设置为全 False 的缓存 GPU 常量张量,确保 + mixed-prefill 逻辑保留这里显式生成的 token,而不会按 decode + 请求重新收集 token。 + 2. 将上一层 model 输出的 hidden states 绑定到 + ``mtp_draft_input_hiddens``,供下一层 MTP model 使用。 + + 参数: + model_input: 当前层的 prefill 输入。方法会原地更新该对象的 + ``input_ids``、``b_is_decode_req`` 和 + ``mtp_draft_input_hiddens`` 字段。 + b_next_token_ids: 每个请求需要追加到当前 query 尾部的 token, + shape 为 ``[batch_size]``。 + mtp_draft_input_hiddens: 上一层 model 产生、供下一层 draft model + 使用的 hidden states。 + + 返回: + 更新后的 ``model_input``,与传入对象是同一个对象。 + + 示例: + 假设 ``b_seq_len=[4, 5]``、``b_ready_cache_len=[1, 2]``,则两个 + 请求当前 query 长度均为 3。若扁平输入和追加 token 为: + + ``input_ids = [10, 11, 12, 20, 21, 22]`` + ``b_next_token_ids = [13, 23]`` + + 更新后的扁平输入为: + + ``input_ids = [11, 12, 13, 21, 22, 23]`` + + 两个请求已缓存的 prefix 长度分别为 1 和 2,它们对应的 token + 不在 ``input_ids`` 中,因此不会被本方法移动或重写。 + """ + model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( + key="dp_vanilla_mtp_prefill_b_is_decode_req", + shape=model_input.b_req_idx.shape, + fill_value=False, + dtype=torch.bool, + ) + model_input.input_ids = gen_mtp_new_input_ids( + input_ids=model_input.input_ids, + b_next_token_ids=b_next_token_ids, + b_seq_len=model_input.b_seq_len, + b_ready_cache_len=model_input.b_ready_cache_len, + ) + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + return model_input diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 945c2334af..3ef0e8c5d1 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -9,12 +9,14 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( BaseSpecProposer, MtpMemIndexesToFree, SpecProposal, ) +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @dataclass @@ -32,9 +34,7 @@ def fill_eagle_draft_model_kv_state( ) -> None: """使用 target prefill 输出初始化 EAGLE draft state。""" - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - prepare_mtp_prefill_inputs( + _prepare_eagle_prefill_inputs( model_input=target_model_input, b_next_token_ids=target_next_token_ids, mtp_draft_input_hiddens=target_model_output.mtp_collector.spec_hidden, @@ -42,6 +42,27 @@ def fill_eagle_draft_model_kv_state( proposer.backend.draft_models[0].forward(target_model_input) +def _prepare_eagle_prefill_inputs( + model_input: ModelInput, + b_next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, +) -> ModelInput: + model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( + key="eagle_mtp_prefill_b_is_decode_req", + shape=model_input.b_req_idx.shape, + fill_value=False, + dtype=torch.bool, + ) + model_input.input_ids = gen_mtp_new_input_ids( + input_ids=model_input.input_ids, + b_next_token_ids=b_next_token_ids, + b_seq_len=model_input.b_seq_len, + b_ready_cache_len=model_input.b_ready_cache_len, + ) + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + return model_input + + def generate_eagle_token_ids( proposer: BaseSpecProposer, model_output: ModelOutput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py index ae6f00b607..d37488d73c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py @@ -17,29 +17,6 @@ class VanillaSpecProposal(SpecProposal): schedule_scores: torch.Tensor | None = None -def fill_chained_mtp_draft_model_kv_state( - proposer: BaseSpecProposer, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, -) -> None: - """构建 Vanilla chained MTP 各级 draft model 的 prefill state。""" - - from lightllm.server.router.model_infer.mode_backend.mtp_pre_process import prepare_mtp_prefill_inputs - - draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_token_ids = target_next_token_ids - for draft_model in proposer.backend.draft_models: - prepare_mtp_prefill_inputs( - model_input=target_model_input, - b_next_token_ids=draft_token_ids, - mtp_draft_input_hiddens=draft_hidden, - ) - draft_output = draft_model.forward(target_model_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - draft_token_ids = proposer.backend._gen_argmax_token_ids(draft_output) - - def propose_next_chained_mtp( proposer: BaseSpecProposer, target_model_input: ModelInput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index 18af6acbef..06c59771e5 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -1,12 +1,15 @@ +import copy + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( VanillaSpecProposal, - fill_chained_mtp_draft_model_kv_state, propose_next_chained_mtp, ) +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager class VanillaWithAttProposer(BaseSpecProposer): @@ -18,12 +21,24 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_chained_mtp_draft_model_kv_state( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - ) + assert target_model_input.is_prefill + assert target_model_input.b_position_delta is None + assert target_next_token_ids.shape == target_model_input.b_req_idx.shape + + # 在局部副本上逐级左移 token 并传递 hidden,避免修改 target + # prefill 输入。Chunked prefill 边界继续用当前级预测 token 补齐。 + draft_input = copy.copy(target_model_input) + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_token_ids = target_next_token_ids + for draft_model in self.backend.draft_models: + draft_input = self._prepare_mtp_prefill_inputs( + model_input=draft_input, + b_next_token_ids=draft_token_ids, + mtp_draft_input_hiddens=draft_hidden, + ) + draft_output = draft_model.forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) def propose_next( self, @@ -43,3 +58,67 @@ def propose_next( draft_step=draft_step, accept_len=accept_len, ) + + def _prepare_mtp_prefill_inputs( + self, + model_input: ModelInput, + b_next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, + ) -> ModelInput: + """构造 Vanilla chained MTP 下一层 draft model 的 prefill 输入。 + + 每个请求当前参与计算的 query 长度为 + ``b_seq_len - b_ready_cache_len``。本方法在各请求自己的 query + 区间内将 token 左移一位,即丢弃区间首 token,并在区间尾部追加该 + 请求的 ``b_next_token_ids``。这样,下一层 MTP model 看到的 token + 与上一层 model 生成的 next token 连续对齐。Chunked prefill 时只 + 移动当前 chunk,已经写入 KV cache 的 prefix 不参与移动。 + + 同时,本方法会完成两项辅助输入设置: + + 1. 将 ``b_is_decode_req`` 设置为全 False 的缓存 GPU 常量张量,确保 + mixed-prefill 逻辑保留这里显式生成的 token,而不会按 decode + 请求重新收集 token。 + 2. 将上一层 model 输出的 hidden states 绑定到 + ``mtp_draft_input_hiddens``,供下一层 MTP model 使用。 + + 参数: + model_input: 当前层的 prefill 输入。方法会原地更新该对象的 + ``input_ids``、``b_is_decode_req`` 和 + ``mtp_draft_input_hiddens`` 字段。 + b_next_token_ids: 每个请求需要追加到当前 query 尾部的 token, + shape 为 ``[batch_size]``。 + mtp_draft_input_hiddens: 上一层 model 产生、供下一层 draft model + 使用的 hidden states。 + + 返回: + 更新后的 ``model_input``,与传入对象是同一个对象。 + + 示例: + 假设 ``b_seq_len=[4, 5]``、``b_ready_cache_len=[1, 2]``,则两个 + 请求当前 query 长度均为 3。若扁平输入和追加 token 为: + + ``input_ids = [10, 11, 12, 20, 21, 22]`` + ``b_next_token_ids = [13, 23]`` + + 更新后的扁平输入为: + + ``input_ids = [11, 12, 13, 21, 22, 23]`` + + 两个请求已缓存的 prefix 长度分别为 1 和 2,它们对应的 token + 不在 ``input_ids`` 中,因此不会被本方法移动或重写。 + """ + model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( + key="vanilla_mtp_prefill_b_is_decode_req", + shape=model_input.b_req_idx.shape, + fill_value=False, + dtype=torch.bool, + ) + model_input.input_ids = gen_mtp_new_input_ids( + input_ids=model_input.input_ids, + b_next_token_ids=b_next_token_ids, + b_seq_len=model_input.b_seq_len, + b_ready_cache_len=model_input.b_ready_cache_len, + ) + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + return model_input diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py index d2ffb793e3..7f0f7a442f 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py @@ -4,6 +4,10 @@ import torch from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_no_att import ( + DpOverlapVanillaNoAttProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_no_att import DpVanillaNoAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer @@ -155,3 +159,17 @@ def test_vanilla_no_att_skips_draft_forward_for_zero_steps(): assert proposal.token_ids.shape == (2, 0) assert proposal.schedule_scores.shape == (2, 0) + + +def test_all_vanilla_no_att_fill_hooks_are_noops(): + backend = SimpleNamespace(draft_models=[]) + proposers = [ + VanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False), + DpVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False), + DpOverlapVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False), + ] + + for proposer in proposers: + proposer.fill_draft_model_kv_state(None, None, None) + + proposers[-1].fill_draft_model_kv_state_overlap(None, None, None, None, None, None) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py new file mode 100644 index 0000000000..d0c0a5157d --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py @@ -0,0 +1,142 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( + fill_dp_chained_mtp_draft_model_kv_state_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_with_att import ( + DpOverlapVanillaWithAttProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import VanillaWithAttProposer + + +def _prefill_input(input_ids, batch_size): + return SimpleNamespace( + is_prefill=True, + b_position_delta=None, + batch_size=batch_size, + input_ids=input_ids, + b_req_idx=torch.arange(batch_size, dtype=torch.int32, device=input_ids.device), + b_seq_len=torch.full( + (batch_size,), input_ids.shape[0] // batch_size, dtype=torch.int32, device=input_ids.device + ), + b_ready_cache_len=torch.zeros(batch_size, dtype=torch.int32, device=input_ids.device), + mtp_draft_input_hiddens=None, + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_chained_prefill_advances_local_input_without_mutating_target(): + device = "cuda" + original_input_ids = torch.tensor([10, 11, 12, 20, 21, 22], dtype=torch.int64, device=device) + target_input = _prefill_input(original_input_ids, batch_size=2) + target_hidden = torch.arange(12, dtype=torch.float32, device=device).reshape(6, 2) + forwarded = [] + + def draft_model(output_tokens, output_hidden): + def forward(model_input): + forwarded.append( + ( + model_input, + model_input.input_ids.clone(), + model_input.mtp_draft_input_hiddens, + model_input.b_is_decode_req, + ) + ) + return SimpleNamespace( + token_ids=output_tokens, + mtp_collector=SimpleNamespace(spec_hidden=output_hidden), + ) + + return SimpleNamespace(forward=forward) + + stage0_hidden = target_hidden + 100 + stage1_hidden = target_hidden + 200 + backend = SimpleNamespace( + draft_models=[ + draft_model(torch.tensor([30, 40], dtype=torch.int64, device=device), stage0_hidden), + draft_model(torch.tensor([31, 41], dtype=torch.int64, device=device), stage1_hidden), + ], + _gen_argmax_token_ids=lambda output: output.token_ids, + ) + proposer = VanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) + + proposer.fill_draft_model_kv_state( + target_model_input=target_input, + target_model_output=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden)), + target_next_token_ids=torch.tensor([13, 23], dtype=torch.int64, device=device), + ) + + assert forwarded[0][0] is forwarded[1][0] + assert forwarded[0][0] is not target_input + torch.testing.assert_close( + forwarded[0][1], + torch.tensor([11, 12, 13, 21, 22, 23], dtype=torch.int64, device=device), + ) + torch.testing.assert_close( + forwarded[1][1], + torch.tensor([12, 13, 30, 22, 23, 40], dtype=torch.int64, device=device), + ) + assert forwarded[0][2] is target_hidden + assert forwarded[1][2] is stage0_hidden + assert not forwarded[0][3].any() + assert forwarded[0][3].data_ptr() == forwarded[1][3].data_ptr() + assert target_input.input_ids is original_input_ids + assert target_input.mtp_draft_input_hiddens is None + + +def test_overlap_chained_prefill_uses_local_microbatch_inputs(monkeypatch): + def prepare(model_input, b_next_token_ids, mtp_draft_input_hiddens): + model_input.input_ids = model_input.input_ids + b_next_token_ids + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + return model_input + + target_input0 = _prefill_input(torch.tensor([1, 2], dtype=torch.int64), batch_size=2) + target_input1 = _prefill_input(torch.tensor([3, 4], dtype=torch.int64), batch_size=2) + target_hidden0 = torch.tensor([[1.0], [2.0]]) + target_hidden1 = torch.tensor([[3.0], [4.0]]) + forwarded = [] + + class DraftModel: + def __init__(self, token_offset): + self.token_offset = token_offset + + def microbatch_overlap_prefill(self, input0, input1): + forwarded.append((input0, input1, input0.input_ids.clone(), input1.input_ids.clone())) + return tuple( + SimpleNamespace( + token_ids=torch.full((2,), self.token_offset + index, dtype=torch.int64), + mtp_collector=SimpleNamespace(spec_hidden=hidden + self.token_offset), + ) + for index, hidden in enumerate((target_hidden0, target_hidden1)) + ) + + backend = SimpleNamespace( + draft_models=[DraftModel(10), DraftModel(20)], + _gen_argmax_token_ids=lambda output: output.token_ids, + ) + proposer = DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) + monkeypatch.setattr(proposer, "_prepare_mtp_prefill_inputs", prepare) + + fill_dp_chained_mtp_draft_model_kv_state_overlap( + proposer=proposer, + target_model_input0=target_input0, + target_model_output0=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden0)), + target_next_token_ids0=torch.tensor([5, 6], dtype=torch.int64), + target_model_input1=target_input1, + target_model_output1=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden1)), + target_next_token_ids1=torch.tensor([7, 8], dtype=torch.int64), + ) + + assert forwarded[0][0] is forwarded[1][0] + assert forwarded[0][1] is forwarded[1][1] + assert forwarded[0][0] is not target_input0 + assert forwarded[0][1] is not target_input1 + torch.testing.assert_close(forwarded[0][2], torch.tensor([6, 8], dtype=torch.int64)) + torch.testing.assert_close(forwarded[0][3], torch.tensor([10, 12], dtype=torch.int64)) + torch.testing.assert_close(forwarded[1][2], torch.tensor([16, 18], dtype=torch.int64)) + torch.testing.assert_close(forwarded[1][3], torch.tensor([21, 23], dtype=torch.int64)) + torch.testing.assert_close(target_input0.input_ids, torch.tensor([1, 2], dtype=torch.int64)) + torch.testing.assert_close(target_input1.input_ids, torch.tensor([3, 4], dtype=torch.int64)) From 2ed52120d480b3ecac7d454b91e40a60913f92bf Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 07:39:46 +0000 Subject: [PATCH 068/103] fix: keep vanilla attention draft depth fixed --- .../mtp_speculative/planner/lightspec.py | 8 ++++++- .../mtp_speculative/test_planner.py | 21 ++++++++++++++++++- 2 files changed, 27 insertions(+), 2 deletions(-) diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index b190285188..24cfbc78b8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -139,7 +139,13 @@ def _get_draft_steps(self) -> Tuple[int, ...]: if self.spec_mode in ("vanilla_no_att", "eagle_no_att"): return tuple(range(self.max_draft_step + 1)) - if self.spec_mode in ("vanilla_with_att", "eagle_with_att", "eagle3"): + # Vanilla with attention uses one draft model for each chained depth. + # Every model only owns the KV state at its fixed cascade position, so + # changing the depth between iterations would leave some levels with + # incomplete or position-misaligned KV. Always run the full chain. + if self.spec_mode == "vanilla_with_att": + return (self.max_draft_step,) + if self.spec_mode in ("eagle_with_att", "eagle3"): return tuple(range(1, self.max_draft_step + 1)) if self.spec_mode == "dflash": return (self.max_draft_step,) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index c4a10cf5df..2187233af8 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -234,7 +234,7 @@ def test_engine_routes_only_dspark_to_the_confidence_planner(): vanilla_planner = build_planner("vanilla_with_att") assert isinstance(vanilla_planner, LightSpecPlanner) assert isinstance(vanilla_planner, BaseMtpPlanner) - assert vanilla_planner.draft_steps == (1, 2, 3) + assert vanilla_planner.draft_steps == (3,) vanilla_no_att_planner = build_planner("vanilla_no_att") assert vanilla_no_att_planner.draft_steps == (0, 1, 2, 3) @@ -480,6 +480,25 @@ def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): assert plan.draft_step == 1 +def test_lightspec_keeps_vanilla_attention_chained_depth_fixed(): + planner = build_lightspec_planner(spec_mode="vanilla_with_att") + for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 10.0)): + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) + for batch_size, draft_cost in ((2, 0.1), (4, 0.1), (8, 0.2)): + planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=draft_cost) + planner._update_verified_batch( + accept_lengths=[4, 4], + req_num=2, + dynamic_batch_size=8, + verified_draft_step=3, + ) + + plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) + + assert planner.draft_steps == (3,) + assert plan.draft_step == plan.pre_draft_step == 3 + + def test_lightspec_compacts_block_verify_without_changing_draft_shape(): planner = build_lightspec_planner( max_draft_step=7, From 515d71dca42e00737fefd49289debaca5b70f566 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 08:16:50 +0000 Subject: [PATCH 069/103] fix: cascade vanilla attention decode inputs --- .../triton_kernel/overlay_mtp_decode_input.py | 74 +++++++++++++ .../proposers/vanilla_with_att.py | 83 ++++++++++++--- .../mtp_speculative/test_vanilla_prefill.py | 100 ++++++++++++++++++ 3 files changed, 245 insertions(+), 12 deletions(-) create mode 100644 lightllm/common/basemodel/triton_kernel/overlay_mtp_decode_input.py diff --git a/lightllm/common/basemodel/triton_kernel/overlay_mtp_decode_input.py b/lightllm/common/basemodel/triton_kernel/overlay_mtp_decode_input.py new file mode 100644 index 0000000000..78a57c004f --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/overlay_mtp_decode_input.py @@ -0,0 +1,74 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _overlay_chained_mtp_decode_input_kernel( + input_ids, + draft_token_ids, + b_req_mtp_start_loc, + accept_len, + BLOCK_SIZE: tl.constexpr, +): + req_index = tl.program_id(0) + req_start = tl.load(b_req_mtp_start_loc + req_index) + req_accept_len = tl.load(accept_len + req_index) + + offsets = tl.arange(0, BLOCK_SIZE) + # tail 行需要保留当前级 draft model 新生成的 token;只覆盖 tail 之前 + # 的已接受行,使它们使用上一层输入中右侧相邻的真实 token。 + mask = offsets < req_accept_len - 1 + shifted_token_ids = tl.load(input_ids + req_start + offsets + 1, mask=mask, other=0) + tl.store(draft_token_ids + req_start + offsets, shifted_token_ids, mask=mask) + + +@torch.no_grad() +def overlay_chained_mtp_decode_input( + input_ids: torch.Tensor, + draft_token_ids: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + accept_len: torch.Tensor, +) -> torch.Tensor: + """将完整 verify token 串与当前 draft token 级联为下一层输入。 + + 对请求 ``i``,令 ``start=b_req_mtp_start_loc[i]``、 + ``tail=start+accept_len[i]-1``。本函数原地执行: + + ``draft_token_ids[start:tail] = input_ids[start + 1:tail + 1]`` + + ``draft_token_ids[tail]`` 保持不变,因此它仍是当前 draft model 在最后 + 接受行生成的新 token。未接受行也保持原 draft 输出,确保所有 token id + 都合法。返回值与传入的 ``draft_token_ids`` 是同一个 tensor。 + + 例如,某请求已接受的输入为 ``[m0, m1, m2]``,当前级在 tail 行生成 + ``d0``,覆盖后该请求的下一层输入就是 ``[m1, m2, d0]``。 + """ + + assert input_ids.is_cuda + assert draft_token_ids.is_cuda + assert b_req_mtp_start_loc.is_cuda + assert accept_len.is_cuda + assert input_ids.shape == draft_token_ids.shape + assert b_req_mtp_start_loc.shape == accept_len.shape + assert input_ids.dtype == draft_token_ids.dtype + assert input_ids.device == draft_token_ids.device == b_req_mtp_start_loc.device == accept_len.device + assert input_ids.is_contiguous() + assert draft_token_ids.is_contiguous() + assert b_req_mtp_start_loc.is_contiguous() + assert accept_len.is_contiguous() + + req_num = int(b_req_mtp_start_loc.shape[0]) + if req_num == 0: + return draft_token_ids + + _overlay_chained_mtp_decode_input_kernel[(req_num,)]( + input_ids=input_ids, + draft_token_ids=draft_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + BLOCK_SIZE=16, + num_warps=1, + num_stages=1, + ) + return draft_token_ids diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index 06c59771e5..82a3a8dc90 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -4,11 +4,9 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids +from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( - VanillaSpecProposal, - propose_next_chained_mtp, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @@ -49,14 +47,75 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=draft_step, - accept_len=accept_len, + """运行完整 Vanilla-With-Att 级联并生成下一轮 proposal。 + + 每一级 draft model 都维护自己所在级联位置的 attention KV,因此 + decode 时必须依次运行全部 draft model,不能动态缩短级联深度。 + 每一级在完整 target verify 布局上 forward,以补齐该级对应位置的 + KV;最终只抽取每个请求最后接受位置的输出作为下一轮候选 token。 + """ + + req_num = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids_by_step = [] + schedule_scores_by_step = [] + + assert not target_model_input.is_prefill + assert accept_len is not None + assert accept_len.shape == (req_num,) + assert target_next_token_ids.shape[0] == target_model_input.batch_size + assert target_model_output.mtp_collector.spec_hidden.shape[0] == target_model_input.batch_size + assert draft_step == self.backend.max_draft_step, ( + "vanilla_with_att requires the full chained draft depth: " + f"draft_step={draft_step}, max_draft_step={self.backend.max_draft_step}" + ) + assert len(self.backend.draft_models) == draft_step + + # target verify 布局中,同一请求的行从 b_req_mtp_start_loc 开始, + # accept_len 包含必然接受的 target token。因此减 1 后得到本轮最后 + # 接受 token 所在的物理行,后续每一级都从这些固定行收集 proposal。 + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + + # draft forward 需要复用 target decode 的请求、位置和 KV slot 布局, + # 但 input_ids 与 hidden 会逐级更新。使用浅副本避免覆盖 target 输入 + # 中这两个字段;布局 tensor 仍然共享,不引入额外拷贝。 + draft_token_ids = target_next_token_ids + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_input = copy.copy(target_model_input) + + for step in range(draft_step): + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_output = self.backend.draft_models[step].forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + + if self.enable_dynmaic_mtp: + # Vanilla With-Att 没有独立的 confidence head;动态调度使用 + # 当前 draft model 选中 token 的采样概率。 + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) + schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) + proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) + + if step + 1 < draft_step: + # 下一层不能直接使用所有行的 draft 预测。已接受前缀继续使用 + # main/上一层输入中的真实 token,仅在 tail 行接上本级新生成 + # 的 draft token,从而形成逐级左移并覆盖尾部的级联输入。 + draft_token_ids = overlay_chained_mtp_decode_input( + input_ids=draft_input.input_ids, + draft_token_ids=draft_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + ) + + proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) + schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) def _prepare_mtp_prefill_inputs( diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py index d0c0a5157d..8a96da3011 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py @@ -140,3 +140,103 @@ def microbatch_overlap_prefill(self, input0, input1): torch.testing.assert_close(forwarded[1][3], torch.tensor([21, 23], dtype=torch.int64)) torch.testing.assert_close(target_input0.input_ids, torch.tensor([1, 2], dtype=torch.int64)) torch.testing.assert_close(target_input1.input_ids, torch.tensor([3, 4], dtype=torch.int64)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_chained_decode_overlays_verified_tokens_for_the_next_level(): + device = "cuda" + original_input_ids = torch.tensor([10, 11, 12, 20, 21, 22], dtype=torch.int64, device=device) + target_input = SimpleNamespace( + is_prefill=False, + batch_size=6, + input_ids=original_input_ids, + mtp_draft_input_hiddens=None, + ) + target_hidden = torch.arange(12, dtype=torch.float32, device=device).reshape(6, 2) + forwarded = [] + + def draft_model(output_tokens, output_probs, output_hidden): + def forward(model_input): + forwarded.append( + ( + model_input, + model_input.input_ids.clone(), + model_input.mtp_draft_input_hiddens, + ) + ) + return SimpleNamespace( + token_ids=output_tokens, + token_probs=output_probs, + mtp_collector=SimpleNamespace(spec_hidden=output_hidden), + ) + + return SimpleNamespace(forward=forward) + + stage0_tokens = torch.tensor([30, 31, 32, 33, 34, 35], dtype=torch.int64, device=device) + stage1_tokens = torch.tensor([40, 41, 42, 43, 44, 45], dtype=torch.int64, device=device) + stage2_tokens = torch.tensor([50, 51, 52, 53, 54, 55], dtype=torch.int64, device=device) + stage0_probs = torch.tensor([0.30, 0.31, 0.32, 0.33, 0.34, 0.35], device=device) + stage1_probs = torch.tensor([0.40, 0.41, 0.42, 0.43, 0.44, 0.45], device=device) + stage2_probs = torch.tensor([0.50, 0.51, 0.52, 0.53, 0.54, 0.55], device=device) + stage0_hidden = target_hidden + 100 + stage1_hidden = target_hidden + 200 + stage2_hidden = target_hidden + 300 + backend = SimpleNamespace( + max_draft_step=3, + draft_models=[ + draft_model(stage0_tokens, stage0_probs, stage0_hidden), + draft_model(stage1_tokens, stage1_probs, stage1_hidden), + draft_model(stage2_tokens, stage2_probs, stage2_hidden), + ], + _gen_argmax_token_ids_and_prob=lambda output: (output.token_ids, output.token_probs), + ) + proposer = VanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=True) + + proposal = proposer.propose_next( + target_model_input=target_input, + target_model_output=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden)), + target_next_token_ids=original_input_ids, + b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32, device=device), + draft_step=3, + accept_len=torch.tensor([3, 2], dtype=torch.int32, device=device), + ) + + assert forwarded[0][0] is forwarded[1][0] is forwarded[2][0] + assert forwarded[0][0] is not target_input + torch.testing.assert_close(forwarded[0][1], original_input_ids) + # 第一次覆盖:[10, 11, 12] -> [11, 12, 32], + # [20, 21] -> [21, 34]。 + torch.testing.assert_close( + forwarded[1][1], + torch.tensor([11, 12, 32, 21, 34, 35], dtype=torch.int64, device=device), + ) + # 第二次覆盖:[11, 12, 32] -> [12, 32, 42], + # [21, 34] -> [34, 44]。 + torch.testing.assert_close( + forwarded[2][1], + torch.tensor([12, 32, 42, 34, 44, 45], dtype=torch.int64, device=device), + ) + assert forwarded[0][2] is target_hidden + assert forwarded[1][2] is stage0_hidden + assert forwarded[2][2] is stage1_hidden + torch.testing.assert_close( + proposal.token_ids, + torch.tensor([[32, 42, 52], [34, 44, 54]], dtype=torch.int64, device=device), + ) + torch.testing.assert_close( + proposal.schedule_scores, + torch.tensor([[0.32, 0.42, 0.52], [0.34, 0.44, 0.54]], device=device), + ) + assert proposal.extra_mem_indexes_cpu == [] + assert target_input.input_ids is original_input_ids + assert target_input.mtp_draft_input_hiddens is None + + with pytest.raises(AssertionError, match="requires the full chained draft depth"): + proposer.propose_next( + target_model_input=target_input, + target_model_output=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden)), + target_next_token_ids=original_input_ids, + b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32, device=device), + draft_step=2, + accept_len=torch.tensor([3, 2], dtype=torch.int32, device=device), + ) From 6fde882a016f04a81786c425f6f0d236d90af6f6 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 08:30:19 +0000 Subject: [PATCH 070/103] refactor: isolate vanilla proposer implementations --- .../dp_overlap_proposers/vanilla_no_att.py | 121 ++++++++++--- .../dp_overlap_proposers/vanilla_utils.py | 107 ----------- .../dp_overlap_proposers/vanilla_with_att.py | 169 ++++++++++++++---- .../dp_proposers/vanilla_no_att.py | 65 +++++-- .../dp_proposers/vanilla_with_att.py | 66 +++++-- .../proposers/vanilla_no_att.py | 11 +- .../proposers/vanilla_utils.py | 63 ------- .../proposers/vanilla_with_att.py | 11 +- .../model_infer/mtp_speculative/utils.py | 10 +- .../mtp_speculative/test_planner.py | 58 +++++- .../mtp_speculative/test_vanilla_overlap.py | 32 ++-- .../mtp_speculative/test_vanilla_prefill.py | 6 +- 12 files changed, 428 insertions(+), 291 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index d70f780f54..f3554d0241 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -1,14 +1,18 @@ +import copy +from dataclasses import dataclass + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( - propose_next_dp_chained_mtp_overlap, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( - VanillaSpecProposal, - propose_next_chained_mtp, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + + +@dataclass +class VanillaSpecProposal(SpecProposal): + """DP-overlap Vanilla No-Att proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None class DpOverlapVanillaNoAttProposer(BaseDpOverlapProposer): @@ -42,14 +46,48 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, + req_num = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids_by_step = [] + schedule_scores_by_step = [] + + if draft_step == 0: + return VanillaSpecProposal( + token_ids=target_next_token_ids.new_empty((req_num, 0)), + extra_mem_indexes_cpu=[], + schedule_scores=( + torch.empty((req_num, 0), dtype=torch.float32, device=target_next_token_ids.device) + if self.enable_dynmaic_mtp + else None + ), + ) + + assert accept_len is not None + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + draft_token_ids = target_next_token_ids + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_input = copy.copy(target_model_input) + + for step in range(draft_step): + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_output = self.backend.draft_models[step].forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) + schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) + proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) + + proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) + schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) def propose_next_overlap( @@ -66,17 +104,44 @@ def propose_next_overlap( accept_len1: torch.Tensor | None, draft_step: int, ) -> VanillaSpecProposal: - return propose_next_dp_chained_mtp_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - real_verify_rows0, - accept_len0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - real_verify_rows1, - accept_len1, - draft_step, + assert accept_len0 is not None + assert accept_len1 is not None + + verify_width = self.backend.max_draft_step + 1 + real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) + req_num_by_batch = tuple(row_count // verify_width for row_count in real_verify_rows) + req_start_rows = ( + torch.arange(0, real_verify_rows0, verify_width, device=target_next_token_ids0.device), + torch.arange(0, real_verify_rows1, verify_width, device=target_next_token_ids1.device), + ) + accepted_tail_rows = ( + req_start_rows[0] + accept_len0[: req_num_by_batch[0]] - 1, + req_start_rows[1] + accept_len1[: req_num_by_batch[1]] - 1, + ) + model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] + draft_token_ids = [target_next_token_ids0, target_next_token_ids1] + draft_hiddens = [ + target_model_output0.mtp_collector.spec_hidden, + target_model_output1.mtp_collector.spec_hidden, + ] + proposal_token_ids = target_next_token_ids0.new_empty((sum(req_num_by_batch), draft_step)) + req_offset = req_num_by_batch[0] + + for step in range(draft_step): + for batch_index, model_input in enumerate(model_inputs): + model_input.input_ids = draft_token_ids[batch_index] + model_input.mtp_draft_input_hiddens = draft_hiddens[batch_index] + + draft_outputs = self.backend.draft_models[step].microbatch_overlap_decode(*model_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden + draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) + + proposal_token_ids[:req_offset, step] = draft_token_ids[0].index_select(0, accepted_tail_rows[0].long()) + proposal_token_ids[req_offset:, step] = draft_token_ids[1].index_select(0, accepted_tail_rows[1].long()) + + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=None, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py deleted file mode 100644 index 328d6a52d8..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_utils.py +++ /dev/null @@ -1,107 +0,0 @@ -"""DP overlap Vanilla proposer 共享辅助函数。""" - -from __future__ import annotations - -import copy - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal - - -def fill_dp_chained_mtp_draft_model_kv_state_overlap( - proposer: BaseDpOverlapProposer, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, -) -> None: - """为两个 DP microbatch 依次构建 Vanilla chained draft state。""" - - target_model_inputs = (target_model_input0, target_model_input1) - target_next_token_ids = (target_next_token_ids0, target_next_token_ids1) - for model_input, next_token_ids in zip(target_model_inputs, target_next_token_ids): - assert model_input.is_prefill - assert model_input.b_position_delta is None - assert next_token_ids.shape == model_input.b_req_idx.shape - - # 两个 microbatch 各自保留逐级左移后的 draft 输入,不修改 - # target prefill ModelInput。 - model_inputs = [copy.copy(model_input) for model_input in target_model_inputs] - draft_hiddens = [ - target_model_output0.mtp_collector.spec_hidden, - target_model_output1.mtp_collector.spec_hidden, - ] - draft_token_ids = list(target_next_token_ids) - - for draft_model in proposer.backend.draft_models: - for batch_index, model_input in enumerate(model_inputs): - model_inputs[batch_index] = proposer._prepare_mtp_prefill_inputs( - model_input=model_input, - b_next_token_ids=draft_token_ids[batch_index], - mtp_draft_input_hiddens=draft_hiddens[batch_index], - ) - draft_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) - for batch_index, draft_output in enumerate(draft_outputs): - draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden - draft_token_ids[batch_index] = proposer.backend._gen_argmax_token_ids(draft_output) - - -def propose_next_dp_chained_mtp_overlap( - proposer: BaseDpOverlapProposer, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, - real_verify_rows1: int, - accept_len1: torch.Tensor, - draft_step: int, -) -> VanillaSpecProposal: - """为两个 DP microbatch 运行 decode,返回按真实请求压缩的 proposal。""" - - model_inputs = (target_model_input0, target_model_input1) - verify_width = proposer.backend.max_draft_step + 1 - real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) - real_request_counts = tuple(row_count // verify_width for row_count in real_verify_rows) - accepted_tail_rows = ( - torch.arange(0, real_verify_rows0, verify_width, device=target_next_token_ids0.device) + accept_len0 - 1, - torch.arange(0, real_verify_rows1, verify_width, device=target_next_token_ids1.device) + accept_len1 - 1, - ) - draft_token_ids = [target_next_token_ids0, target_next_token_ids1] - draft_hiddens = [ - target_model_output0.mtp_collector.spec_hidden, - target_model_output1.mtp_collector.spec_hidden, - ] - proposal_token_ids = target_next_token_ids0.new_empty((sum(real_request_counts), draft_step)) - request_offset = real_request_counts[0] - - for step in range(draft_step): - for batch_index, model_input in enumerate(model_inputs): - model_input.input_ids = draft_token_ids[batch_index] - model_input.mtp_draft_input_hiddens = draft_hiddens[batch_index] - - draft_outputs = proposer.backend.draft_models[step].microbatch_overlap_decode(*model_inputs) - for batch_index, draft_output in enumerate(draft_outputs): - draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden - draft_token_ids[batch_index] = proposer.backend._gen_argmax_token_ids(draft_output) - - proposal_token_ids[:request_offset, step] = draft_token_ids[0].index_select( - 0, accepted_tail_rows[0][:request_offset].long() - ) - proposal_token_ids[request_offset:, step] = draft_token_ids[1].index_select( - 0, accepted_tail_rows[1][: real_request_counts[1]].long() - ) - - return VanillaSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=None, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 2396db8187..9ffad9c674 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -1,21 +1,23 @@ import copy +from dataclasses import dataclass import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids +from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( - fill_dp_chained_mtp_draft_model_kv_state_overlap, - propose_next_dp_chained_mtp_overlap, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( - VanillaSpecProposal, - propose_next_chained_mtp, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +@dataclass +class VanillaSpecProposal(SpecProposal): + """DP-overlap Vanilla With-Att proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + class DpOverlapVanillaWithAttProposer(BaseDpOverlapProposer): """DP ``vanilla_with_att`` proposer。""" @@ -51,15 +53,31 @@ def fill_draft_model_kv_state_overlap( target_model_output1: ModelOutput, target_next_token_ids1: torch.Tensor, ) -> None: - fill_dp_chained_mtp_draft_model_kv_state_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - ) + target_model_inputs = (target_model_input0, target_model_input1) + target_next_token_ids = (target_next_token_ids0, target_next_token_ids1) + for model_input, next_token_ids in zip(target_model_inputs, target_next_token_ids): + assert model_input.is_prefill + assert model_input.b_position_delta is None + assert next_token_ids.shape == model_input.b_req_idx.shape + + model_inputs = [copy.copy(model_input) for model_input in target_model_inputs] + draft_hiddens = [ + target_model_output0.mtp_collector.spec_hidden, + target_model_output1.mtp_collector.spec_hidden, + ] + draft_token_ids = list(target_next_token_ids) + + for draft_model in self.backend.draft_models: + for batch_index, model_input in enumerate(model_inputs): + model_inputs[batch_index] = self._prepare_mtp_prefill_inputs( + model_input=model_input, + b_next_token_ids=draft_token_ids[batch_index], + mtp_draft_input_hiddens=draft_hiddens[batch_index], + ) + draft_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden + draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) def propose_next( self, @@ -70,14 +88,50 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, + req_num = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids_by_step = [] + schedule_scores_by_step = [] + + assert not target_model_input.is_prefill + assert accept_len is not None + assert accept_len.shape == (req_num,) + assert draft_step == self.backend.max_draft_step + assert len(self.backend.draft_models) == draft_step + + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + draft_token_ids = target_next_token_ids + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_input = copy.copy(target_model_input) + + for step in range(draft_step): + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_output = self.backend.draft_models[step].forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) + schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) + proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) + + if step + 1 < draft_step: + draft_token_ids = overlay_chained_mtp_decode_input( + input_ids=draft_input.input_ids, + draft_token_ids=draft_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + ) + + proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) + schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) def propose_next_overlap( @@ -94,19 +148,58 @@ def propose_next_overlap( accept_len1: torch.Tensor | None, draft_step: int, ) -> VanillaSpecProposal: - return propose_next_dp_chained_mtp_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - real_verify_rows0, - accept_len0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - real_verify_rows1, - accept_len1, - draft_step, + assert accept_len0 is not None + assert accept_len1 is not None + assert draft_step == self.backend.max_draft_step + assert len(self.backend.draft_models) == draft_step + + verify_width = self.backend.max_draft_step + 1 + real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) + req_num_by_batch = tuple(row_count // verify_width for row_count in real_verify_rows) + req_start_rows = ( + torch.arange(0, real_verify_rows0, verify_width, device=target_next_token_ids0.device), + torch.arange(0, real_verify_rows1, verify_width, device=target_next_token_ids1.device), + ) + accept_len_by_batch = ( + accept_len0[: req_num_by_batch[0]], + accept_len1[: req_num_by_batch[1]], + ) + accepted_tail_rows = tuple(starts + lengths - 1 for starts, lengths in zip(req_start_rows, accept_len_by_batch)) + model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] + draft_token_ids = [target_next_token_ids0, target_next_token_ids1] + draft_hiddens = [ + target_model_output0.mtp_collector.spec_hidden, + target_model_output1.mtp_collector.spec_hidden, + ] + proposal_token_ids = target_next_token_ids0.new_empty((sum(req_num_by_batch), draft_step)) + req_offset = req_num_by_batch[0] + + for step in range(draft_step): + for batch_index, model_input in enumerate(model_inputs): + model_input.input_ids = draft_token_ids[batch_index] + model_input.mtp_draft_input_hiddens = draft_hiddens[batch_index] + + draft_outputs = self.backend.draft_models[step].microbatch_overlap_decode(*model_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden + draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) + + proposal_token_ids[:req_offset, step] = draft_token_ids[0].index_select(0, accepted_tail_rows[0].long()) + proposal_token_ids[req_offset:, step] = draft_token_ids[1].index_select(0, accepted_tail_rows[1].long()) + + if step + 1 < draft_step: + for batch_index, model_input in enumerate(model_inputs): + draft_token_ids[batch_index] = overlay_chained_mtp_decode_input( + input_ids=model_input.input_ids, + draft_token_ids=draft_token_ids[batch_index], + b_req_mtp_start_loc=req_start_rows[batch_index], + accept_len=accept_len_by_batch[batch_index], + ) + + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=None, ) def _prepare_mtp_prefill_inputs( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py index c7d4ec380c..2249653a2c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -1,11 +1,18 @@ +import copy +from dataclasses import dataclass + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( - VanillaSpecProposal, - propose_next_chained_mtp, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + + +@dataclass +class VanillaSpecProposal(SpecProposal): + """DP Vanilla No-Att proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None class DpVanillaNoAttProposer(BaseDpProposer): @@ -28,12 +35,46 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, + req_num = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids_by_step = [] + schedule_scores_by_step = [] + + if draft_step == 0: + return VanillaSpecProposal( + token_ids=target_next_token_ids.new_empty((req_num, 0)), + extra_mem_indexes_cpu=[], + schedule_scores=( + torch.empty((req_num, 0), dtype=torch.float32, device=target_next_token_ids.device) + if self.enable_dynmaic_mtp + else None + ), + ) + + assert accept_len is not None + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + draft_token_ids = target_next_token_ids + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_input = copy.copy(target_model_input) + + for step in range(draft_step): + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_output = self.backend.draft_models[step].forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) + schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) + proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) + + proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) + schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py index 9c150a4580..b7b4f40714 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -1,17 +1,23 @@ import copy +from dataclasses import dataclass import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids +from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import ( - VanillaSpecProposal, - propose_next_chained_mtp, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +@dataclass +class VanillaSpecProposal(SpecProposal): + """DP Vanilla With-Att proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + class DpVanillaWithAttProposer(BaseDpProposer): """普通 DP ``vanilla_with_att`` proposer。""" @@ -47,14 +53,50 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> VanillaSpecProposal: - return propose_next_chained_mtp( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, + req_num = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids_by_step = [] + schedule_scores_by_step = [] + + assert not target_model_input.is_prefill + assert accept_len is not None + assert accept_len.shape == (req_num,) + assert draft_step == self.backend.max_draft_step + assert len(self.backend.draft_models) == draft_step + + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + draft_token_ids = target_next_token_ids + draft_hidden = target_model_output.mtp_collector.spec_hidden + draft_input = copy.copy(target_model_input) + + for step in range(draft_step): + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_output = self.backend.draft_models[step].forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) + schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) + proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) + + if step + 1 < draft_step: + draft_token_ids = overlay_chained_mtp_decode_input( + input_ids=draft_input.input_ids, + draft_token_ids=draft_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + ) + + proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) + schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) def _prepare_mtp_prefill_inputs( diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index 396720b834..ef7259654b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -1,11 +1,18 @@ import copy +from dataclasses import dataclass import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal + + +@dataclass +class VanillaSpecProposal(SpecProposal): + """Vanilla No-Att proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None class VanillaNoAttProposer(BaseSpecProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py deleted file mode 100644 index d37488d73c..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_utils.py +++ /dev/null @@ -1,63 +0,0 @@ -"""Vanilla chained MTP proposer 共享辅助函数。""" - -from __future__ import annotations - -from dataclasses import dataclass - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal - - -@dataclass -class VanillaSpecProposal(SpecProposal): - """Vanilla proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None - - -def propose_next_chained_mtp( - proposer: BaseSpecProposer, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor, -) -> VanillaSpecProposal: - """依次运行 Vanilla chained MTP 模块并生成 proposal。""" - - request_count = int(b_req_mtp_start_loc.shape[0]) - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_token_ids = target_next_token_ids - draft_hidden = target_model_output.mtp_collector.spec_hidden - proposal_token_ids = target_next_token_ids.new_empty((request_count, draft_step)) - schedule_scores = ( - torch.empty( - (request_count, draft_step), - dtype=torch.float32, - device=target_next_token_ids.device, - ) - if proposer.enable_dynmaic_mtp - else None - ) - - for step in range(draft_step): - draft_model = proposer.backend.draft_models[step] - target_model_input.input_ids = draft_token_ids - target_model_input.mtp_draft_input_hiddens = draft_hidden - draft_output = draft_model.forward(target_model_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - if proposer.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = proposer.backend._gen_argmax_token_ids_and_prob(draft_output) - schedule_scores[:, step] = draft_token_probs.index_select(0, accepted_tail_rows) - else: - draft_token_ids = proposer.backend._gen_argmax_token_ids(draft_output) - proposal_token_ids[:, step] = draft_token_ids.index_select(0, accepted_tail_rows) - - return VanillaSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=schedule_scores, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index 82a3a8dc90..dbea5a797f 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -1,15 +1,22 @@ import copy +from dataclasses import dataclass import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +@dataclass +class VanillaSpecProposal(SpecProposal): + """Vanilla With-Att proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + class VanillaWithAttProposer(BaseSpecProposer): """使用 attention KV cache 的 Vanilla chained MTP proposer。""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index 7933266918..d0b17274ff 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -72,13 +72,9 @@ def scatter_mtp_next_tokens( ) -> None: """Persist the next MTP proposal and optional scheduling scores by request.""" - from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashSpecProposal - from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkSpecProposal - from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import EagleSpecProposal - from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal - - scored_proposal_types = (VanillaSpecProposal, EagleSpecProposal, DFlashSpecProposal, DSparkSpecProposal) - schedule_scores = proposal.schedule_scores if isinstance(proposal, scored_proposal_types) else None + schedule_scores = getattr(proposal, "schedule_scores", None) + if schedule_scores is not None and 0 in schedule_scores.shape: + schedule_scores = None sampling_params_manager = backend.model.req_manager.req_sampling_params_manager mtp_scatter_next_token_ids( diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 2187233af8..d6693490eb 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -61,9 +61,14 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import EagleNoAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import EagleSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_utils import VanillaSpecProposal -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import VanillaWithAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import ( + VanillaNoAttProposer, + VanillaSpecProposal as VanillaNoAttSpecProposal, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import ( + VanillaSpecProposal as VanillaWithAttSpecProposal, + VanillaWithAttProposer, +) from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils @@ -149,9 +154,16 @@ def test_mode_proposals_own_their_schedule_metadata(): assert "schedule_scores" not in SpecProposal.__dataclass_fields__ assert "schedule_scores_cpu" not in SpecProposal.__dataclass_fields__ - for proposal_type in (VanillaSpecProposal, EagleSpecProposal, DFlashSpecProposal, DSparkSpecProposal): + for proposal_type in ( + VanillaNoAttSpecProposal, + VanillaWithAttSpecProposal, + EagleSpecProposal, + DFlashSpecProposal, + DSparkSpecProposal, + ): assert "schedule_scores" in proposal_type.__dataclass_fields__ - assert "schedule_scores_cpu" not in VanillaSpecProposal.__dataclass_fields__ + assert "schedule_scores_cpu" not in VanillaNoAttSpecProposal.__dataclass_fields__ + assert "schedule_scores_cpu" not in VanillaWithAttSpecProposal.__dataclass_fields__ assert "schedule_scores_cpu" not in EagleSpecProposal.__dataclass_fields__ assert "schedule_scores_cpu" not in DFlashSpecProposal.__dataclass_fields__ assert "schedule_scores_cpu" in DSparkSpecProposal.__dataclass_fields__ @@ -199,6 +211,42 @@ def test_scatter_mtp_next_tokens_consumes_mode_proposal(monkeypatch): assert torch.equal(scatter_args["schedule_scores"], proposal.schedule_scores) +def test_scatter_mtp_next_tokens_ignores_empty_schedule_scores(monkeypatch): + scatter_args = {} + monkeypatch.setattr( + mtp_utils, + "mtp_scatter_next_token_ids", + lambda **kwargs: scatter_args.update(kwargs), + ) + backend = SimpleNamespace( + model=SimpleNamespace( + req_manager=SimpleNamespace( + req_sampling_params_manager=SimpleNamespace( + req_to_next_token_ids=torch.empty((4, 2), dtype=torch.int64), + req_to_next_token_scores=torch.empty((4, 2), dtype=torch.float32), + ) + ) + ) + ) + proposal = VanillaNoAttSpecProposal( + token_ids=torch.empty((2, 0), dtype=torch.int64), + extra_mem_indexes_cpu=[], + schedule_scores=torch.empty((2, 0), dtype=torch.float32), + ) + + mtp_utils.scatter_mtp_next_tokens( + backend=backend, + proposal=proposal, + target_next_token_ids=torch.tensor([10, 11], dtype=torch.int64), + b_req_mtp_start_loc=torch.tensor([0, 1], dtype=torch.int32), + b_req_idx=torch.tensor([0, 1], dtype=torch.int32), + mtp_accept_len=torch.ones(2, dtype=torch.int32), + ) + + assert scatter_args["req_to_next_token_scores"] is None + assert scatter_args["schedule_scores"] is None + + def test_dp_planner_returns_fixed_backend_draft_step(): planner = build_dp_planner(backend=SimpleNamespace(max_draft_step=4)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py index 60885ed0e2..307ef1f0e9 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -1,5 +1,6 @@ from types import SimpleNamespace +import pytest import torch from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput @@ -16,14 +17,25 @@ def microbatch_overlap_decode(self, input0, input1): self.decode_batch_sizes.append((input0.batch_size, input1.batch_size)) return tuple( ModelOutput( - logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), - mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), + logits=torch.arange( + model_input.batch_size, + dtype=torch.float32, + device=model_input.input_ids.device, + ).view(-1, 1), + mtp_collector=ModelMtpOutputCollector( + spec_hidden=torch.ones( + (model_input.batch_size, 2), + device=model_input.input_ids.device, + ) + ), ) for model_input in (input0, input1) ) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_dp_vanilla_proposer_owns_overlap_decode(): + device = "cuda" draft_models = [_DraftModel(), _DraftModel()] backend = SimpleNamespace( max_draft_step=2, @@ -37,20 +49,20 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): proposal = proposer.propose_next_overlap( target_model_input0=model_input0, target_model_output0=ModelOutput( - logits=torch.empty((6, 1)), - mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), + logits=torch.empty((6, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), ), - target_next_token_ids0=torch.tensor([10, 11, 0, 0, 0, 0], dtype=torch.int64), + target_next_token_ids0=torch.tensor([10, 11, 0, 0, 0, 0], dtype=torch.int64, device=device), real_verify_rows0=3, - accept_len0=torch.tensor([2], dtype=torch.int32), + accept_len0=torch.tensor([2], dtype=torch.int32, device=device), target_model_input1=model_input1, target_model_output1=ModelOutput( - logits=torch.empty((6, 1)), - mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), + logits=torch.empty((6, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), ), - target_next_token_ids1=torch.tensor([20, 21, 22, 0, 0, 0], dtype=torch.int64), + target_next_token_ids1=torch.tensor([20, 21, 22, 0, 0, 0], dtype=torch.int64, device=device), real_verify_rows1=3, - accept_len1=torch.tensor([1], dtype=torch.int32), + accept_len1=torch.tensor([1], dtype=torch.int32, device=device), draft_step=2, ) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py index 8a96da3011..040edacc03 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py @@ -3,9 +3,6 @@ import pytest import torch -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_utils import ( - fill_dp_chained_mtp_draft_model_kv_state_overlap, -) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_with_att import ( DpOverlapVanillaWithAttProposer, ) @@ -120,8 +117,7 @@ def microbatch_overlap_prefill(self, input0, input1): proposer = DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) monkeypatch.setattr(proposer, "_prepare_mtp_prefill_inputs", prepare) - fill_dp_chained_mtp_draft_model_kv_state_overlap( - proposer=proposer, + proposer.fill_draft_model_kv_state_overlap( target_model_input0=target_input0, target_model_output0=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden0)), target_next_token_ids0=torch.tensor([5, 6], dtype=torch.int64), From d0def7669f16712b394e591d8c28dcf03789fb95 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 08:40:25 +0000 Subject: [PATCH 071/103] refactor: specialize eagle no-att proposal flow --- .../mtp_speculative/proposers/eagle_no_att.py | 103 +++++++++++++++--- .../mtp_speculative/test_eagle_no_att.py | 98 +++++++++++++++++ .../mtp_speculative/test_planner.py | 7 +- 3 files changed, 193 insertions(+), 15 deletions(-) create mode 100644 unit_tests/server/router/model_infer/mtp_speculative/test_eagle_no_att.py diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py index 7c7d12f0f0..e13c499ecc 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -1,14 +1,25 @@ +import copy +from dataclasses import dataclass + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, - fill_eagle_draft_model_kv_state, - propose_next_eagle, +from lightllm.common.basemodel.triton_kernel.select_mtp_rows import ( + select_accepted_tail_rows, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + BaseSpecProposer, + SpecProposal, ) +@dataclass +class EagleSpecProposal(SpecProposal): + """EAGLE No-Att proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + class EagleNoAttProposer(BaseSpecProposer): """不使用 attention KV cache 的 EAGLE proposer。""" @@ -18,7 +29,7 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) + pass def propose_next( self, @@ -29,13 +40,79 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: - return propose_next_eagle( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, + req_num = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids_by_step = [] + schedule_scores_by_step = [] + + if draft_step == 0: + return EagleSpecProposal( + token_ids=target_next_token_ids.new_empty((req_num, 0)), + extra_mem_indexes_cpu=[], + schedule_scores=( + torch.empty( + (req_num, 0), + dtype=torch.float32, + device=target_next_token_ids.device, + ) + if self.enable_dynmaic_mtp + else None + ), + ) + + assert accept_len is not None + # EAGLE No-Att 不维护 draft KV cache。target verification 的输出布局为 + # [verify_batch_size, ...],而下一轮 proposal 每个请求只需要最后一个 + # 接受位置的 token 和 hidden。这里先将所有相关输入压缩到 [req_num, ...], + # 后续各级递归均只计算这 req_num 行。 + selected_rows = select_accepted_tail_rows( b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=draft_step, accept_len=accept_len, - map_draft_token_ids=lambda token_ids: token_ids, + input_ids=target_next_token_ids, + hidden=target_model_output.mtp_collector.spec_hidden, + b_req_idx=target_model_input.b_req_idx, + b_mtp_index=target_model_input.b_mtp_index, + b_seq_len=target_model_input.b_seq_len, + mem_indexes=target_model_input.mem_indexes, + b_shared_seq_len=target_model_input.b_shared_seq_len, + b_shared_radix_node_id=target_model_input.b_shared_radix_node_id, + b_position_delta=target_model_input.b_position_delta, + ) + draft_token_ids = selected_rows.input_ids + draft_hidden = selected_rows.hidden + draft_input = copy.copy(target_model_input) + draft_input.batch_size = req_num + draft_input.input_ids = draft_token_ids + draft_input.b_req_idx = selected_rows.b_req_idx + draft_input.b_mtp_index = selected_rows.b_mtp_index + draft_input.b_seq_len = selected_rows.b_seq_len + draft_input.mem_indexes = selected_rows.mem_indexes + draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len + draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id + draft_input.b_position_delta = selected_rows.b_position_delta + draft_input.mem_indexes_cpu = None + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(req_num)] + + # EAGLE 使用同一个 draft model 递归生成多个 token。每一级将上一级 + # 生成的 token 与 hidden 作为下一次输入;No-Att 模式没有 KV 状态, + # 因此无需分配临时 mem indexes 或修改序列长度。 + draft_model = self.backend.draft_models[0] + for _ in range(draft_step): + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_output = draft_model.forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + schedule_scores_by_step.append(draft_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + proposal_token_ids_by_step.append(draft_token_ids.unsqueeze(1)) + + proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) + schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None + return EagleSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_no_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_no_att.py new file mode 100644 index 0000000000..9cb533e4e8 --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_no_att.py @@ -0,0 +1,98 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import ( + EagleNoAttProposer, +) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_eagle_no_att_recurrently_proposes_from_accepted_tails(): + device = "cuda" + draft_calls = [] + draft_outputs = [ + SimpleNamespace( + token_ids=torch.tensor([21, 24], device=device), + token_probs=torch.tensor([0.8, 0.7], device=device), + mtp_collector=SimpleNamespace(spec_hidden=torch.tensor([[102.0, 103.0], [108.0, 109.0]], device=device)), + ), + SimpleNamespace( + token_ids=torch.tensor([31, 34], device=device), + token_probs=torch.tensor([0.6, 0.5], device=device), + mtp_collector=SimpleNamespace(spec_hidden=torch.tensor([[202.0, 203.0], [208.0, 209.0]], device=device)), + ), + ] + + def forward(model_input): + draft_calls.append( + { + "batch_size": model_input.batch_size, + "input_ids": model_input.input_ids.cpu(), + "draft_hidden": model_input.mtp_draft_input_hiddens.cpu(), + "b_req_idx": model_input.b_req_idx.cpu(), + "b_seq_len": model_input.b_seq_len.cpu(), + "mem_indexes": model_input.mem_indexes.cpu(), + } + ) + return draft_outputs[len(draft_calls) - 1] + + backend = SimpleNamespace( + draft_models=[SimpleNamespace(forward=forward)], + _gen_argmax_token_ids=lambda output: output.token_ids, + _gen_argmax_token_ids_and_prob=lambda output: ( + output.token_ids, + output.token_probs, + ), + ) + proposer = EagleNoAttProposer(backend=backend, enable_dynmaic_mtp=True) + target_model_input = SimpleNamespace( + batch_size=5, + input_ids=torch.tensor([10, 11, 12, 13, 14], device=device), + b_req_idx=torch.tensor([7, 7, 7, 9, 9], dtype=torch.int32, device=device), + b_mtp_index=torch.tensor([0, 1, 2, 0, 1], dtype=torch.int32, device=device), + b_seq_len=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32, device=device), + mem_indexes=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32, device=device), + mem_indexes_cpu=torch.tensor([100, 101, 102, 103, 104], dtype=torch.int32), + b_position_delta=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32, device=device), + b_shared_seq_len=torch.tensor([8, 8, 8, 6, 6], dtype=torch.int32, device=device), + b_shared_radix_node_id=torch.tensor([70, 70, 70, 90, 90], dtype=torch.int64, device=device), + multimodal_params=[{"images": [], "audios": []} for _ in range(5)], + ) + target_model_output = SimpleNamespace( + mtp_collector=SimpleNamespace(spec_hidden=torch.arange(10, dtype=torch.float32, device=device).reshape(5, 2)) + ) + + proposal = proposer.propose_next( + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=target_model_input.input_ids, + b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32, device=device), + draft_step=2, + accept_len=torch.tensor([2, 2], dtype=torch.int32, device=device), + ) + + torch.testing.assert_close(proposal.token_ids, torch.tensor([[21, 31], [24, 34]], device=device)) + torch.testing.assert_close(proposal.schedule_scores, torch.tensor([[0.8, 0.6], [0.7, 0.5]], device=device)) + assert proposal.extra_mem_indexes_cpu == [] + assert len(draft_calls) == 2 + assert draft_calls[0]["batch_size"] == 2 + torch.testing.assert_close(draft_calls[0]["input_ids"], torch.tensor([11, 14])) + torch.testing.assert_close(draft_calls[0]["draft_hidden"], torch.tensor([[2.0, 3.0], [8.0, 9.0]])) + torch.testing.assert_close(draft_calls[0]["b_req_idx"], torch.tensor([7, 9], dtype=torch.int32)) + torch.testing.assert_close(draft_calls[0]["b_seq_len"], torch.tensor([11, 21], dtype=torch.int32)) + torch.testing.assert_close(draft_calls[0]["mem_indexes"], torch.tensor([101, 104], dtype=torch.int32)) + torch.testing.assert_close(draft_calls[1]["input_ids"], torch.tensor([21, 24])) + torch.testing.assert_close( + draft_calls[1]["draft_hidden"], + torch.tensor([[102.0, 103.0], [108.0, 109.0]]), + ) + assert target_model_input.batch_size == 5 + torch.testing.assert_close(target_model_input.input_ids.cpu(), torch.tensor([10, 11, 12, 13, 14])) + + +def test_eagle_no_att_fill_hook_is_noop(): + proposer = EagleNoAttProposer(backend=SimpleNamespace(draft_models=[]), enable_dynmaic_mtp=False) + + proposer.fill_draft_model_kv_state(None, None, None) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index d6693490eb..174c88520b 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -58,7 +58,10 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer, DFlashSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer, DSparkSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import EagleNoAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import ( + EagleNoAttProposer, + EagleSpecProposal as EagleNoAttSpecProposal, +) from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import EagleSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import ( @@ -414,7 +417,7 @@ def test_eagle_proposer_skips_draft_forward_for_zero_steps(): draft_step=0, ) - assert isinstance(proposal, EagleSpecProposal) + assert isinstance(proposal, EagleNoAttSpecProposal) assert proposal.token_ids.shape == (2, 0) assert proposal.schedule_scores.shape == (2, 0) From 5059c4d4f332719e80cdfb345ec1fba73888b580 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 10:20:21 +0000 Subject: [PATCH 072/103] refactor: specialize eagle attention proposal flow --- .../proposers/eagle_with_att.py | 210 ++++++++++++++++-- .../mtp_speculative/test_eagle_with_att.py | 178 +++++++++++++++ 2 files changed, 375 insertions(+), 13 deletions(-) create mode 100644 unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index d5bab30594..7a901220af 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -1,12 +1,25 @@ +import copy +from dataclasses import dataclass + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, - fill_eagle_draft_model_kv_state, - propose_next_eagle, +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids +from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + BaseSpecProposer, + MtpMemIndexesToFree, + SpecProposal, ) +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + + +@dataclass +class EagleSpecProposal(SpecProposal): + """EAGLE With-Att proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None class EagleWithAttProposer(BaseSpecProposer): @@ -18,7 +31,21 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) + assert target_model_input.is_prefill + assert target_model_input.b_position_delta is None + assert target_next_token_ids.shape == target_model_input.b_req_idx.shape + assert len(self.backend.draft_models) == 1 + + # EAGLE 只有一个递归复用的 draft model。用 target prefill 的 token、 + # next token 和 hidden 构造左移一位的 draft 输入,将 prompt KV 写入 + # draft cache。使用浅副本避免覆盖后续流程仍可能读取的 target 输入。 + draft_input = copy.copy(target_model_input) + self._prepare_eagle_prefill_inputs( + model_input=draft_input, + b_next_token_ids=target_next_token_ids, + mtp_draft_input_hiddens=target_model_output.mtp_collector.spec_hidden, + ) + self.backend.draft_models[0].forward(draft_input) def propose_next( self, @@ -29,13 +56,170 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> EagleSpecProposal: - return propose_next_eagle( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, + """提交验证结果对应的 draft KV,并递归生成下一轮 EAGLE proposal。 + + 第一次 draft forward 复用完整 target verify 布局,把本轮所有验证行 + 写入 draft KV;同时从每个请求最后接受的位置生成第一个候选 token。 + 后续级别只保留每个请求一行,使用临时 KV slot 递归生成剩余候选。 + """ + + req_num = int(b_req_mtp_start_loc.shape[0]) + proposal_token_ids_by_step = [] + schedule_scores_by_step = [] + + assert draft_step > 0, "eagle_with_att requires draft_step to be greater than 0 to maintain draft KV state" + assert not target_model_input.is_prefill + assert accept_len is not None + assert accept_len.shape == (req_num,) + assert target_next_token_ids.shape[0] == target_model_input.batch_size + assert target_model_output.mtp_collector.spec_hidden.shape[0] == target_model_input.batch_size + assert target_model_input.b_position_delta is not None + assert len(self.backend.draft_models) == 1 + + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() + draft_model = self.backend.draft_models[0] + + # target verify 的 mem_indexes 正是已验证 token 应写入的 KV slot。 + # 仅替换 token 和 hidden,其他布局 tensor 与 target 共享;浅副本保证 + # draft forward 不会改变 target_model_input 对象自身的字段。 + verify_draft_input = copy.copy(target_model_input) + verify_draft_input.input_ids = target_next_token_ids + verify_draft_input.mtp_draft_input_hiddens = target_model_output.mtp_collector.spec_hidden + extend_output = draft_model.forward(verify_draft_input) + + # 只在 req_num 行 logits 上进行 argmax,避免为未接受的 verify 行执行 + # vocabulary reduction。第一列 proposal 来自每个请求的 accepted tail。 + accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(accepted_tail_output) + schedule_scores_by_step.append(draft_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(accepted_tail_output) + proposal_token_ids_by_step.append(draft_token_ids.unsqueeze(1)) + + if draft_step == 1: + return EagleSpecProposal( + token_ids=torch.cat(proposal_token_ids_by_step, dim=1), + extra_mem_indexes_cpu=[], + schedule_scores=torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None, + ) + + # 后续递归每步、每请求各写一个临时 KV。proposal 在 verify 完成后 + # 统一释放这些 slot,因此同时保留 CPU 索引用于资源回收。 + extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(req_num * (draft_step - 1)) + extra_mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) + + # 一次 Triton kernel 合并抽取 accepted-tail 的 hidden、请求索引、 + # 序列长度、position delta 和共享 radix 元数据,构造 req_num 行的 + # 单 token decode 输入。通用算子同时返回 accepted-tail input ids, + # 因而这里传入与 verify 布局一致的 target_next_token_ids;EAGLE + # 递归 token 来自上面的 draft logits,不使用 selected_rows.input_ids。 + selected_rows = select_accepted_tail_rows( b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=draft_step, accept_len=accept_len, - map_draft_token_ids=lambda token_ids: token_ids, + input_ids=target_next_token_ids, + hidden=extend_output.mtp_collector.spec_hidden, + b_req_idx=target_model_input.b_req_idx, + b_mtp_index=target_model_input.b_mtp_index, + b_seq_len=target_model_input.b_seq_len, + mem_indexes=target_model_input.mem_indexes, + b_shared_seq_len=target_model_input.b_shared_seq_len, + b_shared_radix_node_id=target_model_input.b_shared_radix_node_id, + b_position_delta=target_model_input.b_position_delta, + ) + draft_hidden = selected_rows.hidden + draft_seq_lens = selected_rows.b_seq_len + 1 + max_kv_seq_len = target_model_input.max_kv_seq_len + draft_input = copy.copy(target_model_input) + draft_input.is_prefill = False + draft_input.batch_size = req_num + draft_input.b_req_idx = selected_rows.b_req_idx + draft_input.b_mtp_index = g_pin_mem_manager.get_const_gpu_tensor( + key="eagle_with_att_decode_b_mtp_index", + shape=selected_rows.b_req_idx.shape, + fill_value=0, + dtype=target_model_input.b_mtp_index.dtype, + ) + draft_input.b_seq_len = draft_seq_lens + draft_input.b_position_delta = selected_rows.b_position_delta + draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len + draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id + draft_input.mem_indexes_cpu = None + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(req_num)] + + for step in range(1, draft_step): + mem_start = (step - 1) * req_num + draft_input.input_ids = draft_token_ids + draft_input.mtp_draft_input_hiddens = draft_hidden + draft_input.mem_indexes = extra_mem_indexes[mem_start : mem_start + req_num] + draft_input.max_kv_seq_len = max_kv_seq_len + step + draft_input.total_token_num = req_num * draft_input.max_kv_seq_len + draft_output = draft_model.forward(draft_input) + draft_hidden = draft_output.mtp_collector.spec_hidden + + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + schedule_scores_by_step.append(draft_token_probs.float().unsqueeze(1)) + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + proposal_token_ids_by_step.append(draft_token_ids.unsqueeze(1)) + draft_seq_lens.add_(1) + + proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) + schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None + return EagleSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], + schedule_scores=schedule_scores, + ) + + def _prepare_eagle_prefill_inputs( + self, + model_input: ModelInput, + b_next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, + ) -> ModelInput: + """构造 EAGLE With-Att draft model 的 prompt prefill 输入。 + + EAGLE 的 pre-layer 同时使用 token embedding 和 target hidden。为了让 + draft token 与 target hidden 在预测位置上对齐,本方法会在每个请求 + 当前 query 区间内将 ``input_ids`` 左移一位:移除区间首 token,并 + 将该请求的 ``b_next_token_ids`` 追加到区间末尾。Chunked prefill + 时,已经位于 KV cache 中的 prefix 不参与移动。 + + 此外会把 ``b_is_decode_req`` 设置为缓存的全 False GPU 张量,使 + mixed-prefill 输入整理逻辑保留这里生成的 token;target model 输出的 + hidden 则通过 ``mtp_draft_input_hiddens`` 传给 EAGLE pre-layer。 + + 参数: + model_input: draft prompt prefill 输入。本方法原地更新其 + ``input_ids``、``b_is_decode_req`` 和 hidden 字段。 + b_next_token_ids: 每个请求追加到当前 query 末尾的 next token, + shape 为 ``[batch_size]``。 + mtp_draft_input_hiddens: target prefill 为当前 query 产生的 hidden, + 第一维与扁平 ``input_ids`` 对齐。 + + 返回: + 更新后的 ``model_input``,与传入对象相同。 + + 示例: + 若两个请求当前 query token 分别为 ``[10, 11, 12]`` 和 + ``[20, 21, 22]``,追加 token 为 ``[13, 23]``,则输入从 + ``[10, 11, 12, 20, 21, 22]`` 变为 + ``[11, 12, 13, 21, 22, 23]``。若请求存在 cached prefix,移动 + 仍只发生在当前 query,cache 中的历史 token 和 KV 均保持不变。 + """ + model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( + key="eagle_with_att_prefill_b_is_decode_req", + shape=model_input.b_req_idx.shape, + fill_value=False, + dtype=torch.bool, + ) + model_input.input_ids = gen_mtp_new_input_ids( + input_ids=model_input.input_ids, + b_next_token_ids=b_next_token_ids, + b_seq_len=model_input.b_seq_len, + b_ready_cache_len=model_input.b_ready_cache_len, ) + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + return model_input diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py new file mode 100644 index 0000000000..e5a9d6c6d0 --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py @@ -0,0 +1,178 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import ( + EagleSpecProposal, + EagleWithAttProposer, +) + + +def test_eagle_with_att_rejects_zero_draft_steps(): + proposer = EagleWithAttProposer( + backend=SimpleNamespace(draft_models=[]), + enable_dynmaic_mtp=False, + ) + + with pytest.raises(AssertionError, match="requires draft_step to be greater than 0"): + proposer.propose_next( + target_model_input=None, + target_model_output=None, + target_next_token_ids=torch.tensor([10]), + b_req_mtp_start_loc=torch.tensor([0], dtype=torch.int32), + draft_step=0, + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_eagle_with_att_prefill_builds_draft_kv_without_mutating_target(): + device = "cuda" + original_input_ids = torch.tensor([10, 11, 12, 20, 21, 22], dtype=torch.int64, device=device) + target_hidden = torch.arange(12, dtype=torch.float32, device=device).reshape(6, 2) + target_input = SimpleNamespace( + is_prefill=True, + batch_size=2, + input_ids=original_input_ids, + b_req_idx=torch.tensor([7, 9], dtype=torch.int32, device=device), + b_seq_len=torch.tensor([3, 3], dtype=torch.int32, device=device), + b_ready_cache_len=torch.zeros(2, dtype=torch.int32, device=device), + b_position_delta=None, + b_is_decode_req=None, + mtp_draft_input_hiddens=None, + ) + forwarded = [] + + def forward(model_input): + forwarded.append( + ( + model_input, + model_input.input_ids.clone(), + model_input.mtp_draft_input_hiddens, + model_input.b_is_decode_req, + ) + ) + return ModelOutput(logits=torch.empty((0,), device=device)) + + proposer = EagleWithAttProposer( + backend=SimpleNamespace(draft_models=[SimpleNamespace(forward=forward)]), + enable_dynmaic_mtp=False, + ) + proposer.fill_draft_model_kv_state( + target_model_input=target_input, + target_model_output=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden)), + target_next_token_ids=torch.tensor([13, 23], dtype=torch.int64, device=device), + ) + + assert len(forwarded) == 1 + assert forwarded[0][0] is not target_input + torch.testing.assert_close( + forwarded[0][1], + torch.tensor([11, 12, 13, 21, 22, 23], dtype=torch.int64, device=device), + ) + assert forwarded[0][2] is target_hidden + assert not forwarded[0][3].any() + assert target_input.input_ids is original_input_ids + assert target_input.b_is_decode_req is None + assert target_input.mtp_draft_input_hiddens is None + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_eagle_with_att_commits_verify_kv_then_recurrently_decodes(monkeypatch): + device = "cuda" + original_input_ids = torch.tensor([10, 11, 12, 20, 21, 22], dtype=torch.int64, device=device) + target_hidden = torch.arange(12, dtype=torch.float32, device=device).reshape(6, 2) + extend_hidden = target_hidden + 100 + decode_hidden = torch.tensor([[201.0, 202.0], [203.0, 204.0]], device=device) + draft_calls = [] + + def forward(model_input): + draft_calls.append( + { + "model_input": model_input, + "batch_size": model_input.batch_size, + "input_ids": model_input.input_ids.clone(), + "draft_hidden": model_input.mtp_draft_input_hiddens.clone(), + "b_req_idx": model_input.b_req_idx.clone(), + "b_mtp_index": model_input.b_mtp_index.clone(), + "b_seq_len": model_input.b_seq_len.clone(), + "mem_indexes": model_input.mem_indexes.clone(), + "max_kv_seq_len": model_input.max_kv_seq_len, + "total_token_num": model_input.total_token_num, + } + ) + if len(draft_calls) == 1: + return ModelOutput( + logits=torch.arange(30, 36, dtype=torch.float32, device=device).unsqueeze(1), + mtp_collector=ModelMtpOutputCollector(spec_hidden=extend_hidden), + ) + return ModelOutput( + logits=torch.tensor([[40.0], [41.0]], device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=decode_hidden), + ) + + backend = SimpleNamespace( + draft_models=[SimpleNamespace(forward=forward)], + _gen_argmax_token_ids=lambda output: output.logits[:, 0].long(), + _gen_argmax_token_ids_and_prob=lambda output: ( + output.logits[:, 0].long(), + output.logits[:, 0] / 100, + ), + ) + proposer = EagleWithAttProposer(backend=backend, enable_dynmaic_mtp=True) + target_input = SimpleNamespace( + is_prefill=False, + batch_size=6, + total_token_num=66, + max_kv_seq_len=22, + input_ids=original_input_ids, + b_req_idx=torch.tensor([7, 7, 7, 9, 9, 9], dtype=torch.int32, device=device), + b_mtp_index=torch.tensor([0, 1, 2, 0, 1, 2], dtype=torch.int32, device=device), + b_seq_len=torch.tensor([10, 11, 12, 20, 21, 22], dtype=torch.int32, device=device), + mem_indexes=torch.tensor([100, 101, 102, 103, 104, 105], dtype=torch.int32, device=device), + mem_indexes_cpu=torch.tensor([100, 101, 102, 103, 104, 105], dtype=torch.int32), + b_position_delta=torch.tensor([0, 1, 2, 3, 4, 5], dtype=torch.int32, device=device), + b_shared_seq_len=torch.tensor([8, 8, 8, 6, 6, 6], dtype=torch.int32, device=device), + b_shared_radix_node_id=torch.tensor([70, 70, 70, 90, 90, 90], dtype=torch.int64, device=device), + multimodal_params=[{"images": [], "audios": []} for _ in range(6)], + mtp_draft_input_hiddens=None, + ) + extra_mem_indexes_cpu = torch.tensor([200, 201], dtype=torch.int32) + monkeypatch.setattr(mtp_utils, "alloc_mem_indexes", lambda token_count: extra_mem_indexes_cpu) + + proposal = proposer.propose_next( + target_model_input=target_input, + target_model_output=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden)), + target_next_token_ids=original_input_ids, + b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32, device=device), + draft_step=2, + accept_len=torch.tensor([3, 2], dtype=torch.int32, device=device), + ) + + assert isinstance(proposal, EagleSpecProposal) + torch.testing.assert_close(proposal.token_ids, torch.tensor([[32, 40], [34, 41]], device=device)) + torch.testing.assert_close(proposal.schedule_scores, torch.tensor([[0.32, 0.40], [0.34, 0.41]], device=device)) + assert len(proposal.extra_mem_indexes_cpu) == 1 + assert proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu is extra_mem_indexes_cpu + assert len(draft_calls) == 2 + assert draft_calls[0]["model_input"] is not target_input + assert draft_calls[0]["batch_size"] == 6 + torch.testing.assert_close(draft_calls[0]["input_ids"], original_input_ids) + torch.testing.assert_close(draft_calls[0]["draft_hidden"], target_hidden) + torch.testing.assert_close(draft_calls[0]["mem_indexes"], target_input.mem_indexes) + assert draft_calls[1]["batch_size"] == 2 + torch.testing.assert_close(draft_calls[1]["input_ids"], torch.tensor([32, 34], device=device)) + torch.testing.assert_close( + draft_calls[1]["draft_hidden"], + extend_hidden.index_select(0, torch.tensor([2, 4], device=device)), + ) + torch.testing.assert_close(draft_calls[1]["b_req_idx"], torch.tensor([7, 9], dtype=torch.int32, device=device)) + torch.testing.assert_close(draft_calls[1]["b_mtp_index"], torch.zeros(2, dtype=torch.int32, device=device)) + torch.testing.assert_close(draft_calls[1]["b_seq_len"], torch.tensor([13, 22], dtype=torch.int32, device=device)) + torch.testing.assert_close(draft_calls[1]["mem_indexes"], extra_mem_indexes_cpu.to(device)) + assert draft_calls[1]["max_kv_seq_len"] == 23 + assert draft_calls[1]["total_token_num"] == 46 + assert target_input.input_ids is original_input_ids + assert target_input.mtp_draft_input_hiddens is None From 8442a1875d9f94b320bd00c80b1de55b1892feec Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 10:35:30 +0000 Subject: [PATCH 073/103] refactor: centralize MTP proposal types --- .../dp_overlap_proposers/eagle3.py | 2 +- .../dp_overlap_proposers/eagle_no_att.py | 2 +- .../dp_overlap_proposers/eagle_utils.py | 2 +- .../dp_overlap_proposers/eagle_with_att.py | 2 +- .../dp_overlap_proposers/vanilla_no_att.py | 10 +--- .../dp_overlap_proposers/vanilla_with_att.py | 10 +--- .../mtp_speculative/dp_proposers/eagle3.py | 2 +- .../dp_proposers/eagle_no_att.py | 2 +- .../dp_proposers/eagle_with_att.py | 2 +- .../dp_proposers/vanilla_no_att.py | 10 +--- .../dp_proposers/vanilla_with_att.py | 10 +--- .../mtp_speculative/proposers/eagle3.py | 49 +++-------------- .../mtp_speculative/proposers/eagle_no_att.py | 14 +---- .../mtp_speculative/proposers/eagle_utils.py | 10 +--- .../proposers/eagle_with_att.py | 30 ++++++----- .../proposers/proposal_type.py | 19 +++++++ .../proposers/vanilla_no_att.py | 11 +--- .../proposers/vanilla_with_att.py | 11 +--- .../mtp_speculative/test_eagle_with_att.py | 54 +++++++++++++++++-- .../mtp_speculative/test_planner.py | 32 +++++------ 20 files changed, 123 insertions(+), 161 deletions(-) create mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index 602aa086e0..027e10aa59 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -7,10 +7,10 @@ propose_next_dp_eagle_autoregressive_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, fill_eagle_draft_model_kv_state, propose_next_eagle, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class DpOverlapEagle3Proposer(BaseDpOverlapProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 784e48aca8..a96fc630f4 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -7,10 +7,10 @@ propose_next_dp_eagle_fixed_layout_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, fill_eagle_draft_model_kv_state, propose_next_eagle, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class DpOverlapEagleNoAttProposer(BaseDpOverlapProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index e086b992f7..f751a182aa 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -11,11 +11,11 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, _prepare_eagle_prefill_inputs, generate_eagle_token_ids, prepare_eagle_verify_decode_input, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal def fill_dp_eagle_draft_model_kv_state_overlap( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 15ef075031..85cc5f7696 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -7,10 +7,10 @@ propose_next_dp_eagle_fixed_layout_overlap, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, fill_eagle_draft_model_kv_state, propose_next_eagle, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class DpOverlapEagleWithAttProposer(BaseDpOverlapProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index f3554d0241..5a9961957e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -1,18 +1,10 @@ import copy -from dataclasses import dataclass import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal - - -@dataclass -class VanillaSpecProposal(SpecProposal): - """DP-overlap Vanilla No-Att proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal class DpOverlapVanillaNoAttProposer(BaseDpOverlapProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 9ffad9c674..c95a335bcc 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -1,5 +1,4 @@ import copy -from dataclasses import dataclass import torch @@ -7,17 +6,10 @@ from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -@dataclass -class VanillaSpecProposal(SpecProposal): - """DP-overlap Vanilla With-Att proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None - - class DpOverlapVanillaWithAttProposer(BaseDpOverlapProposer): """DP ``vanilla_with_att`` proposer。""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py index 951357110e..c0b09ae61d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py @@ -3,10 +3,10 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, fill_eagle_draft_model_kv_state, propose_next_eagle, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class DpEagle3Proposer(BaseDpProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py index 052431f287..7ead5c1706 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py @@ -3,10 +3,10 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, fill_eagle_draft_model_kv_state, propose_next_eagle, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class DpEagleNoAttProposer(BaseDpProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py index 53be483e49..abd9f1347c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py @@ -3,10 +3,10 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, fill_eagle_draft_model_kv_state, propose_next_eagle, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class DpEagleWithAttProposer(BaseDpProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py index 2249653a2c..8827ba51d0 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py @@ -1,18 +1,10 @@ import copy -from dataclasses import dataclass import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal - - -@dataclass -class VanillaSpecProposal(SpecProposal): - """DP Vanilla No-Att proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal class DpVanillaNoAttProposer(BaseDpProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py index b7b4f40714..30820c0ee7 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py @@ -1,5 +1,4 @@ import copy -from dataclasses import dataclass import torch @@ -7,17 +6,10 @@ from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -@dataclass -class VanillaSpecProposal(SpecProposal): - """DP Vanilla With-Att proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None - - class DpVanillaWithAttProposer(BaseDpProposer): """普通 DP ``vanilla_with_att`` proposer。""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py index 544a762f7e..4ce867ff16 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py @@ -1,52 +1,19 @@ import torch -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( - EagleSpecProposal, - fill_eagle_draft_model_kv_state, - generate_eagle_token_ids, - generate_eagle_token_ids_and_prob, - propose_next_eagle, -) +from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer -class Eagle3Proposer(BaseSpecProposer): - """带有 draft-to-target vocabulary 映射的 EAGLE3 proposer。""" +class Eagle3Proposer(EagleWithAttProposer): + """复用 EAGLE With-Att KV 流程并执行 draft-to-target 词表映射。""" def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: - return generate_eagle_token_ids(self, model_output, self._map_draft_token_ids) + draft_token_ids = super()._gen_argmax_token_ids(model_output) + return self._map_draft_token_ids(draft_token_ids) def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): - return generate_eagle_token_ids_and_prob(self, model_output, self._map_draft_token_ids) - - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) - - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> EagleSpecProposal: - return propose_next_eagle( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=draft_step, - accept_len=accept_len, - map_draft_token_ids=self._map_draft_token_ids, - ) + draft_token_ids, draft_token_probs = super()._gen_argmax_token_ids_and_prob(model_output) + return self._map_draft_token_ids(draft_token_ids), draft_token_probs diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py index e13c499ecc..629eb7fd3e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -1,5 +1,4 @@ import copy -from dataclasses import dataclass import torch @@ -7,17 +6,8 @@ from lightllm.common.basemodel.triton_kernel.select_mtp_rows import ( select_accepted_tail_rows, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( - BaseSpecProposer, - SpecProposal, -) - - -@dataclass -class EagleSpecProposal(SpecProposal): - """EAGLE No-Att proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class EagleNoAttProposer(BaseSpecProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py index 3ef0e8c5d1..bece44064d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py @@ -3,7 +3,6 @@ from __future__ import annotations import copy -from dataclasses import dataclass from typing import Callable import torch @@ -14,18 +13,11 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( BaseSpecProposer, MtpMemIndexesToFree, - SpecProposal, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -@dataclass -class EagleSpecProposal(SpecProposal): - """EAGLE proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None - - def fill_eagle_draft_model_kv_state( proposer: BaseSpecProposer, target_model_input: ModelInput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index 7a901220af..3d2c0a0e86 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -1,5 +1,4 @@ import copy -from dataclasses import dataclass import torch @@ -10,18 +9,11 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( BaseSpecProposer, MtpMemIndexesToFree, - SpecProposal, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -@dataclass -class EagleSpecProposal(SpecProposal): - """EAGLE With-Att proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None - - class EagleWithAttProposer(BaseSpecProposer): """使用 attention KV cache 的 EAGLE proposer。""" @@ -67,7 +59,7 @@ def propose_next( proposal_token_ids_by_step = [] schedule_scores_by_step = [] - assert draft_step > 0, "eagle_with_att requires draft_step to be greater than 0 to maintain draft KV state" + assert draft_step > 0, "EAGLE attention requires draft_step to be greater than 0 to maintain draft KV state" assert not target_model_input.is_prefill assert accept_len is not None assert accept_len.shape == (req_num,) @@ -91,10 +83,10 @@ def propose_next( # vocabulary reduction。第一列 proposal 来自每个请求的 accepted tail。 accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) if self.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(accepted_tail_output) + draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(accepted_tail_output) schedule_scores_by_step.append(draft_token_probs.float().unsqueeze(1)) else: - draft_token_ids = self.backend._gen_argmax_token_ids(accepted_tail_output) + draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) proposal_token_ids_by_step.append(draft_token_ids.unsqueeze(1)) if draft_step == 1: @@ -158,10 +150,10 @@ def propose_next( draft_hidden = draft_output.mtp_collector.spec_hidden if self.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(draft_output) schedule_scores_by_step.append(draft_token_probs.float().unsqueeze(1)) else: - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + draft_token_ids = self._gen_argmax_token_ids(draft_output) proposal_token_ids_by_step.append(draft_token_ids.unsqueeze(1)) draft_seq_lens.add_(1) @@ -173,6 +165,16 @@ def propose_next( schedule_scores=schedule_scores, ) + def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: + """生成 target vocabulary 下的候选 token;EAGLE3 会覆盖词表映射。""" + + return self.backend._gen_argmax_token_ids(model_output) + + def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): + """生成候选 token 及其概率;EAGLE3 会覆盖 token 的词表映射。""" + + return self.backend._gen_argmax_token_ids_and_prob(model_output) + def _prepare_eagle_prefill_inputs( self, model_input: ModelInput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py new file mode 100644 index 0000000000..2d9531936c --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py @@ -0,0 +1,19 @@ +from dataclasses import dataclass + +import torch + +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + + +@dataclass +class EagleSpecProposal(SpecProposal): + """EAGLE proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + +@dataclass +class VanillaSpecProposal(SpecProposal): + """Vanilla proposal with optional selected-token probabilities.""" + + schedule_scores: torch.Tensor | None = None diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py index ef7259654b..ea6a2d5e11 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -1,18 +1,11 @@ import copy -from dataclasses import dataclass import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal - - -@dataclass -class VanillaSpecProposal(SpecProposal): - """Vanilla No-Att proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal class VanillaNoAttProposer(BaseSpecProposer): diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index dbea5a797f..87a8376848 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -1,22 +1,15 @@ import copy -from dataclasses import dataclass import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -@dataclass -class VanillaSpecProposal(SpecProposal): - """Vanilla With-Att proposal with optional selected-token probabilities.""" - - schedule_scores: torch.Tensor | None = None - - class VanillaWithAttProposer(BaseSpecProposer): """使用 attention KV cache 的 Vanilla chained MTP proposer。""" diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py index e5a9d6c6d0..f50780d53c 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py @@ -5,10 +5,56 @@ from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import ( - EagleSpecProposal, - EagleWithAttProposer, -) +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal + + +def test_eagle3_reuses_attention_flow_and_maps_proposal_tokens(): + target_input_ids = torch.tensor([10, 20], dtype=torch.int64) + target_hidden = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) + + def forward(model_input): + assert model_input.input_ids is target_input_ids + assert model_input.mtp_draft_input_hiddens is target_hidden + return ModelOutput( + logits=torch.tensor([[1.0], [2.0]]), + mtp_collector=ModelMtpOutputCollector(spec_hidden=target_hidden + 10), + ) + + draft_model = SimpleNamespace( + forward=forward, + map_draft_vocab_to_main_vocab=lambda token_ids: token_ids + 100, + ) + proposer = Eagle3Proposer( + backend=SimpleNamespace( + draft_models=[draft_model], + _gen_argmax_token_ids=lambda output: output.logits[:, 0].long(), + ), + enable_dynmaic_mtp=False, + ) + target_input = SimpleNamespace( + is_prefill=False, + batch_size=2, + input_ids=None, + b_position_delta=torch.zeros(2, dtype=torch.int32), + mtp_draft_input_hiddens=None, + ) + + proposal = proposer.propose_next( + target_model_input=target_input, + target_model_output=SimpleNamespace(mtp_collector=SimpleNamespace(spec_hidden=target_hidden)), + target_next_token_ids=target_input_ids, + b_req_mtp_start_loc=torch.tensor([0, 1], dtype=torch.int32), + draft_step=1, + accept_len=torch.ones(2, dtype=torch.int32), + ) + + assert isinstance(proposal, EagleSpecProposal) + torch.testing.assert_close(proposal.token_ids, torch.tensor([[101], [102]])) + assert proposal.schedule_scores is None + assert target_input.input_ids is None + assert target_input.mtp_draft_input_hiddens is None def test_eagle_with_att_rejects_zero_draft_steps(): diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 174c88520b..77c60e4118 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -58,20 +58,14 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer, DFlashSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer, DSparkSpecProposal from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import ( - EagleNoAttProposer, - EagleSpecProposal as EagleNoAttSpecProposal, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import EagleSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import EagleNoAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import ( - VanillaNoAttProposer, - VanillaSpecProposal as VanillaNoAttSpecProposal, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import ( - VanillaSpecProposal as VanillaWithAttSpecProposal, - VanillaWithAttProposer, +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + EagleSpecProposal, + VanillaSpecProposal, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import VanillaWithAttProposer from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils @@ -158,15 +152,13 @@ def test_mode_proposals_own_their_schedule_metadata(): assert "schedule_scores_cpu" not in SpecProposal.__dataclass_fields__ for proposal_type in ( - VanillaNoAttSpecProposal, - VanillaWithAttSpecProposal, + VanillaSpecProposal, EagleSpecProposal, DFlashSpecProposal, DSparkSpecProposal, ): assert "schedule_scores" in proposal_type.__dataclass_fields__ - assert "schedule_scores_cpu" not in VanillaNoAttSpecProposal.__dataclass_fields__ - assert "schedule_scores_cpu" not in VanillaWithAttSpecProposal.__dataclass_fields__ + assert "schedule_scores_cpu" not in VanillaSpecProposal.__dataclass_fields__ assert "schedule_scores_cpu" not in EagleSpecProposal.__dataclass_fields__ assert "schedule_scores_cpu" not in DFlashSpecProposal.__dataclass_fields__ assert "schedule_scores_cpu" in DSparkSpecProposal.__dataclass_fields__ @@ -231,7 +223,7 @@ def test_scatter_mtp_next_tokens_ignores_empty_schedule_scores(monkeypatch): ) ) ) - proposal = VanillaNoAttSpecProposal( + proposal = VanillaSpecProposal( token_ids=torch.empty((2, 0), dtype=torch.int64), extra_mem_indexes_cpu=[], schedule_scores=torch.empty((2, 0), dtype=torch.float32), @@ -320,13 +312,12 @@ def test_dynamic_planner_registers_cuda_graph_costs_from_backend(): assert planner.draft_infer_costs.estimate(4) == 0.3 -def test_each_mode_proposer_only_inherits_its_base_interface(): +def test_each_mode_proposer_inherits_its_expected_implementation_base(): proposer_types = ( VanillaWithAttProposer, VanillaNoAttProposer, EagleWithAttProposer, EagleNoAttProposer, - Eagle3Proposer, DFlashProposer, DSparkProposer, ) @@ -347,6 +338,7 @@ def test_each_mode_proposer_only_inherits_its_base_interface(): for proposer_type in proposer_types: assert proposer_type.__bases__ == (BaseSpecProposer,) + assert Eagle3Proposer.__bases__ == (EagleWithAttProposer,) for proposer_type in dp_proposer_types: assert proposer_type.__bases__ == (BaseDpProposer,) for proposer_type in dp_overlap_proposer_types: @@ -417,7 +409,7 @@ def test_eagle_proposer_skips_draft_forward_for_zero_steps(): draft_step=0, ) - assert isinstance(proposal, EagleNoAttSpecProposal) + assert isinstance(proposal, EagleSpecProposal) assert proposal.token_ids.shape == (2, 0) assert proposal.schedule_scores.shape == (2, 0) From 320d7f990712550fabcea80b097fe2206824201f Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 10:40:52 +0000 Subject: [PATCH 074/103] refactor: relocate DP EAGLE helpers --- .../mtp_speculative/dp_overlap_proposers/eagle3.py | 4 +--- .../mtp_speculative/dp_overlap_proposers/eagle_no_att.py | 4 +--- .../mtp_speculative/dp_overlap_proposers/eagle_utils.py | 6 ++++-- .../mtp_speculative/dp_overlap_proposers/eagle_with_att.py | 4 +--- .../model_infer/mtp_speculative/dp_proposers/eagle3.py | 2 +- .../mtp_speculative/dp_proposers/eagle_no_att.py | 2 +- .../{proposers => dp_proposers}/eagle_utils.py | 2 +- .../mtp_speculative/dp_proposers/eagle_with_att.py | 2 +- 8 files changed, 11 insertions(+), 15 deletions(-) rename lightllm/server/router/model_infer/mtp_speculative/{proposers => dp_proposers}/eagle_utils.py (99%) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index 027e10aa59..9e947893ab 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -4,10 +4,8 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, - propose_next_dp_eagle_autoregressive_overlap, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( fill_eagle_draft_model_kv_state, + propose_next_dp_eagle_autoregressive_overlap, propose_next_eagle, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index a96fc630f4..8383f71ec1 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -4,10 +4,8 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, - propose_next_dp_eagle_fixed_layout_overlap, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( fill_eagle_draft_model_kv_state, + propose_next_dp_eagle_fixed_layout_overlap, propose_next_eagle, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index f751a182aa..e276540f67 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -9,12 +9,14 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_utils import ( _prepare_eagle_prefill_inputs, + fill_eagle_draft_model_kv_state, generate_eagle_token_ids, prepare_eagle_verify_decode_input, + propose_next_eagle, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 85cc5f7696..efb179c11a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -4,10 +4,8 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, - propose_next_dp_eagle_fixed_layout_overlap, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( fill_eagle_draft_model_kv_state, + propose_next_dp_eagle_fixed_layout_overlap, propose_next_eagle, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py index c0b09ae61d..77e2a8e354 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py @@ -2,7 +2,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_utils import ( fill_eagle_draft_model_kv_state, propose_next_eagle, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py index 7ead5c1706..16479336d9 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py @@ -2,7 +2,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_utils import ( fill_eagle_draft_model_kv_state, propose_next_eagle, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_utils.py similarity index 99% rename from lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py rename to lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_utils.py index bece44064d..c647137904 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_utils.py @@ -1,4 +1,4 @@ -"""EAGLE proposer 共享辅助函数。""" +"""普通 DP 与 DP-overlap EAGLE proposer 共用的辅助函数。""" from __future__ import annotations diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py index abd9f1347c..fa36357eda 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py @@ -2,7 +2,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_utils import ( +from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_utils import ( fill_eagle_draft_model_kv_state, propose_next_eagle, ) From 6cd7d8abb5b06c579a5477df13e9485d3ef22668 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 12:04:05 +0000 Subject: [PATCH 075/103] refactor: specialize parallel MTP proposers --- .../layer_infer/pre_layer_infer.py | 17 +-- lightllm/models/qwen3_eagle/model.py | 31 +++- .../mtp_speculative/planner/dspark.py | 2 +- .../mtp_speculative/proposers/dflash.py | 139 +++++++++++++----- .../mtp_speculative/proposers/dspark.py | 139 ++++++++++++++---- .../proposers/parallel_block_utils.py | 101 ------------- .../proposers/proposal_type.py | 15 ++ .../models/test_qwen3_eagle_model_input.py | 81 ++++++++++ .../mtp_speculative/test_dflash.py | 125 ++++++++++++++++ .../mtp_speculative/test_dspark.py | 138 +++++++++++++++++ .../mtp_speculative/test_planner.py | 6 +- unit_tests/utils/test_speculative_utils.py | 128 ---------------- 12 files changed, 614 insertions(+), 308 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py create mode 100644 unit_tests/models/test_qwen3_eagle_model_input.py create mode 100644 unit_tests/server/router/model_infer/mtp_speculative/test_dflash.py create mode 100644 unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py diff --git a/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py index f437a7d28c..80290cfae4 100644 --- a/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py +++ b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py @@ -4,7 +4,7 @@ class Qwen3EaglePreLayerInfer(LlamaPreLayerInfer): - """Eagle3 draft-token embedding plus target-hidden projection.""" + """EAGLE3 draft-token embedding and fixed-width draft hidden preparation.""" def __init__(self, network_config): super().__init__(network_config) @@ -13,14 +13,11 @@ def __init__(self, network_config): def prepare_spec_draft_hiddens( self, infer_state: InferStateInfo, - layer_weight: Qwen3EaglePreAndPostLayerWeight, ) -> None: - target_hiddens = infer_state.mtp_draft_input_hiddens - # Target verification provides concatenated auxiliary-layer hiddens (N * H). - # Autoregressive draft steps feed the previous draft output, which is already H. - if target_hiddens.shape[-1] != self.hidden_size_: - target_hiddens = layer_weight.fc_weight_.mm(target_hiddens) - infer_state.eagle_draft_hidden_states = target_hiddens + draft_hiddens = infer_state.mtp_draft_input_hiddens + assert draft_hiddens is not None + assert draft_hiddens.shape[-1] == self.hidden_size_ + infer_state.eagle_draft_hidden_states = draft_hiddens def context_forward( self, @@ -28,7 +25,7 @@ def context_forward( infer_state: InferStateInfo, layer_weight: Qwen3EaglePreAndPostLayerWeight, ): - self.prepare_spec_draft_hiddens(infer_state, layer_weight) + self.prepare_spec_draft_hiddens(infer_state) return super().context_forward(input_ids, infer_state, layer_weight) def token_forward( @@ -37,5 +34,5 @@ def token_forward( infer_state: InferStateInfo, layer_weight: Qwen3EaglePreAndPostLayerWeight, ): - self.prepare_spec_draft_hiddens(infer_state, layer_weight) + self.prepare_spec_draft_hiddens(infer_state) return super().token_forward(input_ids, infer_state, layer_weight) diff --git a/lightllm/models/qwen3_eagle/model.py b/lightllm/models/qwen3_eagle/model.py index 83c90ebbde..8ab63705aa 100644 --- a/lightllm/models/qwen3_eagle/model.py +++ b/lightllm/models/qwen3_eagle/model.py @@ -1,6 +1,9 @@ -import torch +import copy from typing import List +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.models.llama.model import LlamaTpPartModel from lightllm.models.draft_registry import DraftModelRegistry @@ -78,6 +81,32 @@ def _init_infer_layer(self, start_layer_index=None): ) super()._init_infer_layer(start_layer_index=total_pre_layers_num) + def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): + """在进入模型主体和 CUDA Graph 前统一 EAGLE3 hidden 的宽度。 + + Target model 收集的多个辅助层 hidden 会沿最后一维拼接,宽度为 + ``target_layer_num * hidden_size``;后续自回归 draft step 返回的 hidden + 已经是 ``hidden_size``。Decode CUDA Graph 只按 batch size 缓存,因此 + 不能让这两种 shape 和对应的 FC 分支进入同一张图。 + + 这里在创建 InferState 之前完成一次 target hidden 投影,使 prefill、 + decode、Graph capture 和 replay 看到的输入始终为 ``[token_num, hidden_size]``。 + 使用浅副本避免覆盖调用方持有的 target ModelInput。 + """ + + draft_hiddens = model_input.mtp_draft_input_hiddens + assert draft_hiddens is not None + assert draft_hiddens.ndim == 2 + assert draft_hiddens.shape[0] == model_input.input_ids.shape[0] + + hidden_size = int(self.config["hidden_size"]) + if draft_hiddens.shape[-1] != hidden_size: + model_input = copy.copy(model_input) + model_input.mtp_draft_input_hiddens = self.pre_post_weight.fc_weight_.mm(draft_hiddens) + + assert model_input.mtp_draft_input_hiddens.shape == (model_input.input_ids.shape[0], hidden_size) + return super()._create_inferstate(model_input=model_input, microbatch_index=microbatch_index) + # d2t stores per-token offsets: target_id = draft_id + d2t[draft_id]. @torch.no_grad() def map_draft_vocab_to_main_vocab(self, draft_token_ids: torch.Tensor) -> torch.Tensor: diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py index 552644bb83..e7a6828d66 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -10,7 +10,7 @@ SpecDecodePlan, _InferCostMsTable, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import DSparkSpecProposal if TYPE_CHECKING: from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py index dff3f67309..23f81f546a 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -1,27 +1,19 @@ from __future__ import annotations -from dataclasses import dataclass +import copy import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( BaseSpecProposer, MtpMemIndexesToFree, - SpecProposal, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( - build_parallel_block_draft_input, - fill_parallel_block_draft_model_kv_state, - extend_parallel_block_draft_kv_cache, +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + DFlashSpecProposal, ) - - -@dataclass -class DFlashSpecProposal(SpecProposal): - """DFlash proposal with optional block-token probabilities.""" - - schedule_scores: torch.Tensor | None = None +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager class DFlashProposer(BaseSpecProposer): @@ -37,12 +29,23 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_parallel_block_draft_model_kv_state( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - ) + assert target_model_input.is_prefill + assert target_model_input.b_position_delta is None + assert target_next_token_ids.shape == target_model_input.b_req_idx.shape + assert len(self.backend.draft_models) == 1 + + target_hidden = target_model_output.mtp_collector.spec_hidden + assert target_hidden is not None + assert target_hidden.shape[0] == target_model_input.input_ids.shape[0] + if target_hidden.numel() == 0: + return + + # DFlash prefill 直接复用 target prompt 的 token 布局,并注入 target + # hidden 初始化唯一 draft model 的 KV。使用浅副本避免在 target 输入上 + # 保留 draft 专用状态。 + draft_input = copy.copy(target_model_input) + draft_input.mtp_draft_input_hiddens = target_hidden + self.backend.draft_models[0].forward(draft_input) @torch.no_grad() def propose_next( @@ -54,36 +57,106 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> DFlashSpecProposal: - request_count = int(b_req_mtp_start_loc.shape[0]) + """提交 target verify KV,并生成下一轮 DFlash block proposal。 + + 首次 draft forward 复用完整 target verify 布局,把本轮所有验证行的 + target hidden 写入 draft KV。随后每个请求以最后接受的 token 作为 + anchor,在其后填充 mask token,并通过一次 non-causal block forward + 并行生成完整候选块。 + """ + + req_num = int(b_req_mtp_start_loc.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) - # One accepted-tail anchor expands to a complete block-diffusion draft. + assert draft_step > 0, "DFlash requires draft_step to be greater than 0 to maintain draft KV state" + assert draft_step <= block_size + assert not target_model_input.is_prefill + assert accept_len is not None + assert accept_len.shape == (req_num,) + assert target_next_token_ids.shape[0] == target_model_input.batch_size + assert target_model_output.mtp_collector.spec_hidden is not None + assert target_model_output.mtp_collector.spec_hidden.shape[0] == target_model_input.batch_size + assert target_model_input.b_position_delta is not None + assert len(self.backend.draft_models) == 1 + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_input, extra_mem_indexes_cpu = build_parallel_block_draft_input( - proposer=self, - target_model_input=target_model_input, - target_next_token_ids=target_next_token_ids, - accepted_tail_rows=accepted_tail_rows, - request_count=request_count, + + # target verify 的行布局和 mem_indexes 对应本轮所有被验证 token。 + # 附加 target hidden 后执行一次 draft forward,将这些行提交到 DFlash + # KV cache;浅副本保证 target_model_input 本身保持不变。 + verify_draft_input = copy.copy(target_model_input) + verify_draft_input.mtp_draft_input_hiddens = target_model_output.mtp_collector.spec_hidden + draft_model.forward(verify_draft_input) + + # 每个请求始终展开完整 block,未被本轮 proposal 返回的 block 尾部仍会 + # 参与 parallel forward。所有临时 KV slot 在 verify 后通过 proposal + # 统一释放。 + extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(req_num * block_size) + block_input_ids = target_next_token_ids.new_full( + (req_num * block_size,), + fill_value=draft_model.mask_token_id, + ) + block_input_ids[::block_size] = target_next_token_ids.index_select(0, accepted_tail_rows) + + block_offsets = torch.arange( + block_size, + dtype=target_model_input.b_seq_len.dtype, + device=target_next_token_ids.device, + ) + draft_input = copy.copy(target_model_input) + draft_input.input_ids = block_input_ids + draft_input.mtp_draft_input_hiddens = None + draft_input.total_token_num = req_num * block_size + draft_input.batch_size = draft_input.total_token_num + draft_input.max_q_seq_len = 1 + draft_input.max_kv_seq_len = target_model_input.max_kv_seq_len + block_size + draft_input.b_req_idx = ( + target_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() + ) + draft_input.b_mtp_index = g_pin_mem_manager.get_const_gpu_tensor( + key="dflash_decode_b_mtp_index", + shape=draft_input.b_req_idx.shape, + fill_value=0, + dtype=target_model_input.b_mtp_index.dtype, + ) + draft_input.b_seq_len = ( + (target_model_input.b_seq_len.index_select(0, accepted_tail_rows)[:, None] + block_offsets[None, :] + 1) + .reshape(-1) + .contiguous() + ) + draft_input.b_position_delta = ( + target_model_input.b_position_delta.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() + ) + draft_input.b_shared_seq_len = ( + target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() ) - extend_parallel_block_draft_kv_cache( - proposer=self, - target_model_input=target_model_input, - target_hidden=target_model_output.mtp_collector.spec_hidden, + draft_input.b_shared_radix_node_id = ( + target_model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() ) + draft_input.mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) + draft_input.mem_indexes_cpu = None + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(draft_input.batch_size)] draft_output = draft_model.forward(draft_input) if self.enable_dynmaic_mtp: flat_draft_token_ids, flat_draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) else: flat_draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) + assert flat_draft_token_ids.numel() == req_num * block_size + block_draft_token_ids = flat_draft_token_ids.reshape(req_num, block_size) proposal_token_ids = block_draft_token_ids[:, :draft_step].contiguous() schedule_scores = None if self.enable_dynmaic_mtp: - block_draft_token_probs = flat_draft_token_probs.reshape(request_count, block_size) + assert flat_draft_token_probs.numel() == req_num * block_size + block_draft_token_probs = flat_draft_token_probs.reshape(req_num, block_size) schedule_scores = block_draft_token_probs[:, :draft_step].float().contiguous() return DFlashSpecProposal( token_ids=proposal_token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index f7d22f18b8..5e3ea5694f 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -1,31 +1,21 @@ from __future__ import annotations -from dataclasses import dataclass +import copy import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( BaseSpecProposer, MtpMemIndexesToFree, - SpecProposal, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( - build_parallel_block_draft_input, - fill_parallel_block_draft_model_kv_state, - extend_parallel_block_draft_kv_cache, +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + DSparkSpecProposal, ) -@dataclass -class DSparkSpecProposal(SpecProposal): - """DSpark proposal with GPU confidence scores and their CPU planner view.""" - - schedule_scores: torch.Tensor | None = None - schedule_scores_cpu: torch.Tensor | None = None - - class DSparkProposer(BaseSpecProposer): """DSpark semi-autoregressive parallel-block proposer. @@ -40,12 +30,23 @@ def fill_draft_model_kv_state( target_model_output: ModelOutput, target_next_token_ids: torch.Tensor, ) -> None: - fill_parallel_block_draft_model_kv_state( - proposer=self, - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - ) + assert target_model_input.is_prefill + assert target_model_input.b_position_delta is None + assert target_next_token_ids.shape == target_model_input.b_req_idx.shape + assert len(self.backend.draft_models) == 1 + + target_hidden = target_model_output.mtp_collector.spec_hidden + assert target_hidden is not None + assert target_hidden.shape[0] == target_model_input.input_ids.shape[0] + if target_hidden.numel() == 0: + return + + # DSpark prefill 直接使用 target prompt 的 token 布局和 hidden,将 prompt + # KV 写入唯一的 parallel-block draft model。使用浅副本,避免把 draft + # 专用 hidden 挂到后续流程仍可能读取的 target ModelInput 上。 + draft_input = copy.copy(target_model_input) + draft_input.mtp_draft_input_hiddens = target_hidden + self.backend.draft_models[0].forward(draft_input) @torch.no_grad() def propose_next( @@ -57,37 +58,111 @@ def propose_next( draft_step: int, accept_len: torch.Tensor | None = None, ) -> DSparkSpecProposal: - request_count = int(b_req_mtp_start_loc.shape[0]) + """提交 target verify KV,并生成下一轮 DSpark block proposal。 + + 首次 draft forward 复用完整 target verify 布局,把本轮所有验证行的 + target hidden 写入 draft KV。随后每个请求取最后接受的 token 作为 + block anchor,并在其后填充 mask token;第二次 forward 一次生成完整 + block。Markov head 可以直接返回具有块内依赖的 token,confidence head + 则为动态 verify 提供每个 draft 位置的条件置信度。 + """ + + req_num = int(b_req_mtp_start_loc.shape[0]) draft_model = self.backend.draft_models[0] block_size = int(draft_model.block_size) schedule_scores = None + assert draft_step > 0, "DSpark requires draft_step to be greater than 0 to maintain draft KV state" + assert draft_step <= block_size + assert not target_model_input.is_prefill + assert accept_len is not None + assert accept_len.shape == (req_num,) + assert target_next_token_ids.shape[0] == target_model_input.batch_size + assert target_model_output.mtp_collector.spec_hidden is not None + assert target_model_output.mtp_collector.spec_hidden.shape[0] == target_model_input.batch_size + assert target_model_input.b_position_delta is not None + assert len(self.backend.draft_models) == 1 + accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_input, extra_mem_indexes_cpu = build_parallel_block_draft_input( - proposer=self, - target_model_input=target_model_input, - target_next_token_ids=target_next_token_ids, - accepted_tail_rows=accepted_tail_rows, - request_count=request_count, + + # target verify 的行布局和 mem_indexes 已经对应本轮所有被验证 token。 + # 仅附加 target hidden 后执行一次 draft forward,即可把这些行提交到 + # DSpark KV cache;浅副本保证 target_model_input 本身不被修改。 + verify_draft_input = copy.copy(target_model_input) + verify_draft_input.mtp_draft_input_hiddens = target_model_output.mtp_collector.spec_hidden + draft_model.forward(verify_draft_input) + + # DSpark 每个请求固定展开一个完整 block,临时 KV 在 target verify 完成 + # 后通过 proposal 统一释放。block 第一行是 accepted-tail anchor,其余行 + # 使用 mask token,由 parallel backbone 一次并行计算。 + extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(req_num * block_size) + block_input_ids = target_next_token_ids.new_full( + (req_num * block_size,), + fill_value=draft_model.mask_token_id, + ) + block_input_ids[::block_size] = target_next_token_ids.index_select(0, accepted_tail_rows) + + block_offsets = torch.arange( + block_size, + dtype=target_model_input.b_seq_len.dtype, + device=target_next_token_ids.device, + ) + draft_input = copy.copy(target_model_input) + draft_input.input_ids = block_input_ids + draft_input.mtp_draft_input_hiddens = None + draft_input.total_token_num = req_num * block_size + draft_input.batch_size = draft_input.total_token_num + draft_input.max_q_seq_len = 1 + draft_input.max_kv_seq_len = target_model_input.max_kv_seq_len + block_size + draft_input.b_req_idx = ( + target_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() + ) + draft_input.b_mtp_index = g_pin_mem_manager.get_const_gpu_tensor( + key="dspark_decode_b_mtp_index", + shape=draft_input.b_req_idx.shape, + fill_value=0, + dtype=target_model_input.b_mtp_index.dtype, + ) + draft_input.b_seq_len = ( + (target_model_input.b_seq_len.index_select(0, accepted_tail_rows)[:, None] + block_offsets[None, :] + 1) + .reshape(-1) + .contiguous() + ) + draft_input.b_position_delta = ( + target_model_input.b_position_delta.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() + ) + draft_input.b_shared_seq_len = ( + target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() ) - extend_parallel_block_draft_kv_cache( - proposer=self, - target_model_input=target_model_input, - target_hidden=target_model_output.mtp_collector.spec_hidden, + draft_input.b_shared_radix_node_id = ( + target_model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) + .repeat_interleave(block_size) + .contiguous() ) + draft_input.mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) + draft_input.mem_indexes_cpu = None + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(draft_input.batch_size)] draft_output = draft_model.forward(draft_input) if draft_output.mtp_collector.draft_token_ids is None: flat_draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) else: flat_draft_token_ids = draft_output.mtp_collector.draft_token_ids - block_draft_token_ids = flat_draft_token_ids.reshape(request_count, block_size) + assert flat_draft_token_ids.numel() == req_num * block_size + block_draft_token_ids = flat_draft_token_ids.reshape(req_num, block_size) proposal_token_ids = block_draft_token_ids[:, :draft_step].contiguous() if self.enable_dynmaic_mtp: confidence_logits = draft_output.mtp_collector.confidence_logits if confidence_logits is None: raise RuntimeError("DSpark dynamic verify requires confidence head logits") + assert confidence_logits.ndim == 2 + assert confidence_logits.shape[0] == req_num + assert confidence_logits.shape[1] >= draft_step # Match the clamp used by the GPU dynamic row selector before it # converts conditional confidence to prefix survival probability. schedule_scores = ( diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py deleted file mode 100644 index 989238ca2a..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/parallel_block_utils.py +++ /dev/null @@ -1,101 +0,0 @@ -"""Parallel-block proposer 共享辅助函数。""" - -from __future__ import annotations - -import copy - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer - - -@torch.no_grad() -def fill_parallel_block_draft_model_kv_state( - proposer: BaseSpecProposer, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, -) -> None: - """使用 target hidden 初始化 parallel-block drafter 的 KV state。""" - - target_hidden = target_model_output.mtp_collector.spec_hidden - if target_hidden.numel() == 0: - return - - target_model_input.mtp_draft_input_hiddens = target_hidden - proposer.backend.draft_models[0].forward(target_model_input) - - -def extend_parallel_block_draft_kv_cache( - proposer: BaseSpecProposer, - target_model_input: ModelInput, - target_hidden: torch.Tensor, -) -> None: - """用 MTP decode 提交本轮 target verify hidden,扩展 drafter KV。""" - - assert not target_model_input.is_prefill - assert target_model_input.b_position_delta is not None - target_model_input.mtp_draft_input_hiddens = target_hidden - proposer.backend.draft_models[0].forward(target_model_input) - - -def build_parallel_block_draft_input( - proposer: BaseSpecProposer, - target_model_input: ModelInput, - target_next_token_ids: torch.Tensor, - accepted_tail_rows: torch.Tensor, - request_count: int, -): - """构建 parallel-block drafter 的 accepted-token + mask-token 输入。""" - - draft_model = proposer.backend.draft_models[0] - block_size = int(draft_model.block_size) - extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(request_count * block_size) - - block_input_ids = target_next_token_ids.new_full( - (request_count * block_size,), - fill_value=draft_model.mask_token_id, - ) - block_input_ids[::block_size] = target_next_token_ids.index_select(0, accepted_tail_rows) - - block_offsets = torch.arange( - block_size, - dtype=target_model_input.b_seq_len.dtype, - device=target_next_token_ids.device, - ) - draft_input = copy.copy(target_model_input) - draft_input.input_ids = block_input_ids - draft_input.total_token_num = draft_input.input_ids.shape[0] - draft_input.batch_size = draft_input.total_token_num - draft_input.max_q_seq_len = 1 - draft_input.max_kv_seq_len = target_model_input.max_kv_seq_len + block_size - draft_input.b_req_idx = ( - target_model_input.b_req_idx.index_select(0, accepted_tail_rows).repeat_interleave(block_size).contiguous() - ) - draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) - draft_input.b_seq_len = ( - (target_model_input.b_seq_len.index_select(0, accepted_tail_rows)[:, None] + block_offsets[None, :] + 1) - .reshape(-1) - .contiguous() - ) - draft_input.b_position_delta = ( - target_model_input.b_position_delta.index_select(0, accepted_tail_rows) - .repeat_interleave(block_size) - .contiguous() - ) - draft_input.b_shared_seq_len = ( - target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) - .repeat_interleave(block_size) - .contiguous() - ) - draft_input.b_shared_radix_node_id = ( - target_model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) - .repeat_interleave(block_size) - .contiguous() - ) - draft_input.mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids.device, non_blocking=True) - empty_multimodal_params = {"images": [], "audios": []} - draft_input.multimodal_params = [empty_multimodal_params] * draft_input.batch_size - return draft_input, extra_mem_indexes_cpu diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py index 2d9531936c..ee70a26616 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py @@ -17,3 +17,18 @@ class VanillaSpecProposal(SpecProposal): """Vanilla proposal with optional selected-token probabilities.""" schedule_scores: torch.Tensor | None = None + + +@dataclass +class DFlashSpecProposal(SpecProposal): + """DFlash proposal with optional block-token probabilities.""" + + schedule_scores: torch.Tensor | None = None + + +@dataclass +class DSparkSpecProposal(SpecProposal): + """DSpark proposal with GPU confidence scores and their CPU planner view.""" + + schedule_scores: torch.Tensor | None = None + schedule_scores_cpu: torch.Tensor | None = None diff --git a/unit_tests/models/test_qwen3_eagle_model_input.py b/unit_tests/models/test_qwen3_eagle_model_input.py new file mode 100644 index 0000000000..5ffd9f73fb --- /dev/null +++ b/unit_tests/models/test_qwen3_eagle_model_input.py @@ -0,0 +1,81 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.qwen3_eagle.layer_infer.pre_layer_infer import Qwen3EaglePreLayerInfer +from lightllm.models.qwen3_eagle.model import Qwen3EagleModel + + +def test_qwen3_eagle_projects_concatenated_target_hiddens_before_inferstate(monkeypatch): + parent_calls = [] + monkeypatch.setattr( + LlamaTpPartModel, + "_create_inferstate", + lambda self, model_input, microbatch_index=0: parent_calls.append((model_input, microbatch_index)) + or SimpleNamespace(model_input=model_input), + ) + projected_hiddens = torch.empty((3, 4)) + projection_inputs = [] + model = Qwen3EagleModel.__new__(Qwen3EagleModel) + model.config = {"hidden_size": 4} + model.pre_post_weight = SimpleNamespace( + fc_weight_=SimpleNamespace( + mm=lambda hiddens: projection_inputs.append(hiddens) or projected_hiddens, + ) + ) + target_hiddens = torch.empty((3, 12)) + model_input = SimpleNamespace( + input_ids=torch.arange(3), + mtp_draft_input_hiddens=target_hiddens, + ) + + infer_state = model._create_inferstate(model_input=model_input, microbatch_index=1) + + assert projection_inputs == [target_hiddens] + assert len(parent_calls) == 1 + normalized_input, microbatch_index = parent_calls[0] + assert normalized_input is not model_input + assert normalized_input.mtp_draft_input_hiddens is projected_hiddens + assert model_input.mtp_draft_input_hiddens is target_hiddens + assert microbatch_index == 1 + assert infer_state.model_input is normalized_input + + +def test_qwen3_eagle_keeps_recursive_draft_hiddens_without_projection(monkeypatch): + monkeypatch.setattr( + LlamaTpPartModel, + "_create_inferstate", + lambda self, model_input, microbatch_index=0: model_input, + ) + model = Qwen3EagleModel.__new__(Qwen3EagleModel) + model.config = {"hidden_size": 4} + model.pre_post_weight = SimpleNamespace( + fc_weight_=SimpleNamespace( + mm=lambda _: pytest.fail("fixed-width recursive hidden must not be projected"), + ) + ) + model_input = SimpleNamespace( + input_ids=torch.arange(3), + mtp_draft_input_hiddens=torch.empty((3, 4)), + ) + + normalized_input = model._create_inferstate(model_input=model_input) + + assert normalized_input is model_input + + +def test_qwen3_eagle_pre_layer_rejects_non_normalized_hidden_width(): + pre_layer = Qwen3EaglePreLayerInfer.__new__(Qwen3EaglePreLayerInfer) + pre_layer.hidden_size_ = 4 + infer_state = SimpleNamespace(mtp_draft_input_hiddens=torch.empty((2, 12))) + + with pytest.raises(AssertionError): + pre_layer.prepare_spec_draft_hiddens(infer_state) + + normalized_hiddens = torch.empty((2, 4)) + infer_state.mtp_draft_input_hiddens = normalized_hiddens + pre_layer.prepare_spec_draft_hiddens(infer_state) + + assert infer_state.eagle_draft_hidden_states is normalized_hiddens diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_dflash.py b/unit_tests/server/router/model_infer/mtp_speculative/test_dflash.py new file mode 100644 index 0000000000..ac3a303075 --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_dflash.py @@ -0,0 +1,125 @@ +from types import SimpleNamespace + +import torch + +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils +from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import DFlashSpecProposal +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + + +def test_dflash_prefill_uses_a_shallow_copy_for_target_hidden(): + forwarded_inputs = [] + draft_model = SimpleNamespace(forward=forwarded_inputs.append) + proposer = DFlashProposer( + backend=SimpleNamespace(draft_models=[draft_model]), + enable_dynmaic_mtp=False, + ) + model_input = SimpleNamespace( + is_prefill=True, + b_position_delta=None, + b_req_idx=torch.tensor([3, 5], dtype=torch.int32), + input_ids=torch.arange(5, dtype=torch.int64), + mtp_draft_input_hiddens=None, + ) + target_hidden = torch.empty((5, 8)) + + proposer.fill_draft_model_kv_state( + target_model_input=model_input, + target_model_output=SimpleNamespace( + mtp_collector=SimpleNamespace(spec_hidden=target_hidden), + ), + target_next_token_ids=torch.tensor([11, 13], dtype=torch.int64), + ) + + assert len(forwarded_inputs) == 1 + assert forwarded_inputs[0] is not model_input + assert forwarded_inputs[0].mtp_draft_input_hiddens is target_hidden + assert model_input.mtp_draft_input_hiddens is None + + +def test_dflash_commits_verify_kv_and_builds_parallel_block(monkeypatch): + block_size = 3 + flat_draft_token_ids = torch.tensor([30, 31, 32, 40, 41, 42], dtype=torch.int64) + flat_draft_token_probs = torch.tensor([0.9, 0.8, 0.7, 0.6, 0.5, 0.4], dtype=torch.float16) + forwarded_inputs = [] + + def forward(model_input): + forwarded_inputs.append(model_input) + return SimpleNamespace() + + draft_model = SimpleNamespace( + block_size=block_size, + mask_token_id=99, + forward=forward, + ) + proposer = DFlashProposer( + backend=SimpleNamespace( + draft_models=[draft_model], + _gen_argmax_token_ids_and_prob=lambda _: (flat_draft_token_ids, flat_draft_token_probs), + ), + enable_dynmaic_mtp=True, + ) + extra_mem_indexes_cpu = torch.arange(100, 106, dtype=torch.int32) + monkeypatch.setattr(mtp_utils, "alloc_mem_indexes", lambda token_num: extra_mem_indexes_cpu) + monkeypatch.setattr(torch.Tensor, "cuda", lambda self, non_blocking=False: self) + monkeypatch.setattr( + g_pin_mem_manager, + "get_const_gpu_tensor", + lambda *, shape, fill_value, dtype, **_: torch.full(shape, fill_value, dtype=dtype), + ) + + target_hidden = torch.empty((5, 8)) + model_input = SimpleNamespace( + is_prefill=False, + batch_size=5, + total_token_num=40, + max_q_seq_len=1, + max_kv_seq_len=9, + input_ids=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int64), + b_req_idx=torch.tensor([7, 7, 7, 9, 9], dtype=torch.int32), + b_mtp_index=torch.arange(5, dtype=torch.int32), + b_seq_len=torch.tensor([4, 5, 6, 8, 9], dtype=torch.int32), + b_position_delta=torch.tensor([1, 1, 1, 2, 2], dtype=torch.int32), + b_shared_seq_len=torch.tensor([3, 3, 3, 6, 6], dtype=torch.int32), + b_shared_radix_node_id=torch.tensor([17, 17, 17, 19, 19], dtype=torch.int64), + mem_indexes=torch.arange(5, dtype=torch.int32), + mem_indexes_cpu=torch.arange(5, dtype=torch.int32), + multimodal_params=[{"images": [], "audios": []} for _ in range(5)], + mtp_draft_input_hiddens=None, + ) + + proposal = proposer.propose_next( + target_model_input=model_input, + target_model_output=SimpleNamespace( + mtp_collector=SimpleNamespace(spec_hidden=target_hidden), + ), + target_next_token_ids=model_input.input_ids, + b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32), + draft_step=2, + accept_len=torch.tensor([2, 2], dtype=torch.int32), + ) + + assert isinstance(proposal, DFlashSpecProposal) + assert len(forwarded_inputs) == 2 + verify_draft_input, block_draft_input = forwarded_inputs + assert verify_draft_input is not model_input + assert verify_draft_input.mtp_draft_input_hiddens is target_hidden + assert model_input.mtp_draft_input_hiddens is None + + assert torch.equal(block_draft_input.input_ids, torch.tensor([11, 99, 99, 21, 99, 99])) + assert torch.equal(block_draft_input.b_req_idx, torch.tensor([7, 7, 7, 9, 9, 9], dtype=torch.int32)) + assert torch.equal(block_draft_input.b_mtp_index, torch.zeros(6, dtype=torch.int32)) + assert torch.equal(block_draft_input.b_seq_len, torch.tensor([6, 7, 8, 10, 11, 12], dtype=torch.int32)) + assert torch.equal(block_draft_input.b_position_delta, torch.tensor([1, 1, 1, 2, 2, 2], dtype=torch.int32)) + assert block_draft_input.mtp_draft_input_hiddens is None + assert block_draft_input.mem_indexes is extra_mem_indexes_cpu + assert block_draft_input.mem_indexes_cpu is None + + assert torch.equal(proposal.token_ids, torch.tensor([[30, 31], [40, 41]], dtype=torch.int64)) + torch.testing.assert_close( + proposal.schedule_scores, + flat_draft_token_probs.reshape(2, block_size)[:, :2].float(), + ) + assert len(proposal.extra_mem_indexes_cpu) == 1 + assert proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu is extra_mem_indexes_cpu diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py b/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py new file mode 100644 index 0000000000..d341dc7dbd --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py @@ -0,0 +1,138 @@ +from types import SimpleNamespace + +import torch + +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils +from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + + +def test_dspark_prefill_uses_a_shallow_copy_for_target_hidden(): + forwarded_inputs = [] + draft_model = SimpleNamespace(forward=forwarded_inputs.append) + proposer = DSparkProposer( + backend=SimpleNamespace(draft_models=[draft_model]), + enable_dynmaic_mtp=False, + ) + model_input = SimpleNamespace( + is_prefill=True, + b_position_delta=None, + b_req_idx=torch.tensor([3, 5], dtype=torch.int32), + input_ids=torch.arange(5, dtype=torch.int64), + mtp_draft_input_hiddens=None, + ) + target_hidden = torch.empty((5, 8)) + + proposer.fill_draft_model_kv_state( + target_model_input=model_input, + target_model_output=SimpleNamespace( + mtp_collector=SimpleNamespace(spec_hidden=target_hidden), + ), + target_next_token_ids=torch.tensor([11, 13], dtype=torch.int64), + ) + + assert len(forwarded_inputs) == 1 + assert forwarded_inputs[0] is not model_input + assert forwarded_inputs[0].mtp_draft_input_hiddens is target_hidden + assert model_input.mtp_draft_input_hiddens is None + + +def test_dspark_commits_verify_kv_and_builds_parallel_block(monkeypatch): + block_size = 3 + draft_token_ids = torch.tensor([30, 31, 32, 40, 41, 42], dtype=torch.int64) + confidence_logits = torch.tensor( + [ + [-10.0, 0.0, 10.0], + [10.0, 0.0, -10.0], + ] + ) + block_output = SimpleNamespace( + mtp_collector=SimpleNamespace( + draft_token_ids=draft_token_ids, + confidence_logits=confidence_logits, + ) + ) + forwarded_inputs = [] + + def forward(model_input): + forwarded_inputs.append(model_input) + return block_output if len(forwarded_inputs) == 2 else SimpleNamespace() + + draft_model = SimpleNamespace( + block_size=block_size, + mask_token_id=99, + forward=forward, + ) + proposer = DSparkProposer( + backend=SimpleNamespace(draft_models=[draft_model]), + enable_dynmaic_mtp=True, + ) + extra_mem_indexes_cpu = torch.arange(100, 106, dtype=torch.int32) + monkeypatch.setattr(mtp_utils, "alloc_mem_indexes", lambda token_num: extra_mem_indexes_cpu) + monkeypatch.setattr(torch.Tensor, "cuda", lambda self, non_blocking=False: self) + monkeypatch.setattr( + g_pin_mem_manager, + "get_const_gpu_tensor", + lambda *, shape, fill_value, dtype, **_: torch.full(shape, fill_value, dtype=dtype), + ) + monkeypatch.setattr( + g_pin_mem_manager, + "async_copy_from_gpu_tensor", + lambda *, gpu_tensor, **_: gpu_tensor.clone(), + ) + + target_hidden = torch.empty((5, 8)) + model_input = SimpleNamespace( + is_prefill=False, + batch_size=5, + total_token_num=40, + max_q_seq_len=1, + max_kv_seq_len=9, + input_ids=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int64), + b_req_idx=torch.tensor([7, 7, 7, 9, 9], dtype=torch.int32), + b_mtp_index=torch.arange(5, dtype=torch.int32), + b_seq_len=torch.tensor([4, 5, 6, 8, 9], dtype=torch.int32), + b_position_delta=torch.tensor([1, 1, 1, 2, 2], dtype=torch.int32), + b_shared_seq_len=torch.tensor([3, 3, 3, 6, 6], dtype=torch.int32), + b_shared_radix_node_id=torch.tensor([17, 17, 17, 19, 19], dtype=torch.int64), + mem_indexes=torch.arange(5, dtype=torch.int32), + mem_indexes_cpu=torch.arange(5, dtype=torch.int32), + multimodal_params=[{"images": [], "audios": []} for _ in range(5)], + mtp_draft_input_hiddens=None, + ) + + proposal = proposer.propose_next( + target_model_input=model_input, + target_model_output=SimpleNamespace( + mtp_collector=SimpleNamespace(spec_hidden=target_hidden), + ), + target_next_token_ids=model_input.input_ids, + b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32), + draft_step=2, + accept_len=torch.tensor([2, 2], dtype=torch.int32), + ) + + assert len(forwarded_inputs) == 2 + verify_draft_input, block_draft_input = forwarded_inputs + assert verify_draft_input is not model_input + assert verify_draft_input.mtp_draft_input_hiddens is target_hidden + assert model_input.mtp_draft_input_hiddens is None + + assert torch.equal(block_draft_input.input_ids, torch.tensor([11, 99, 99, 21, 99, 99])) + assert torch.equal(block_draft_input.b_req_idx, torch.tensor([7, 7, 7, 9, 9, 9], dtype=torch.int32)) + assert torch.equal(block_draft_input.b_mtp_index, torch.zeros(6, dtype=torch.int32)) + assert torch.equal(block_draft_input.b_seq_len, torch.tensor([6, 7, 8, 10, 11, 12], dtype=torch.int32)) + assert torch.equal(block_draft_input.b_position_delta, torch.tensor([1, 1, 1, 2, 2, 2], dtype=torch.int32)) + assert block_draft_input.mtp_draft_input_hiddens is None + assert block_draft_input.mem_indexes is extra_mem_indexes_cpu + assert block_draft_input.mem_indexes_cpu is None + + assert torch.equal(proposal.token_ids, torch.tensor([[30, 31], [40, 41]], dtype=torch.int64)) + torch.testing.assert_close( + proposal.schedule_scores, + torch.tensor([[0.01, 0.5], [0.99, 0.5]], dtype=torch.float32), + ) + assert torch.equal(proposal.schedule_scores_cpu, proposal.schedule_scores) + assert proposal.schedule_scores_cpu is not proposal.schedule_scores + assert len(proposal.extra_mem_indexes_cpu) == 1 + assert proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu is extra_mem_indexes_cpu diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 77c60e4118..737412d408 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -55,12 +55,14 @@ MtpMemIndexesToFree, SpecProposal, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer, DFlashSpecProposal -from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer, DSparkSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import EagleNoAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + DFlashSpecProposal, + DSparkSpecProposal, EagleSpecProposal, VanillaSpecProposal, ) diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py index ba2afdfd09..aaa3e05268 100644 --- a/unit_tests/utils/test_speculative_utils.py +++ b/unit_tests/utils/test_speculative_utils.py @@ -13,13 +13,6 @@ from lightllm.common.basemodel.attention.triton.fp import TritonDecodeAttState from lightllm.models import get_draft_model_class from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight -from lightllm.server.router.model_infer.mtp_speculative.proposers import dflash as dflash_module -from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer, DFlashSpecProposal -from lightllm.server.router.model_infer.mtp_speculative.proposers.parallel_block_utils import ( - build_parallel_block_draft_input, - extend_parallel_block_draft_kv_cache, -) -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.utils import envs_utils @@ -266,127 +259,6 @@ def test_draft_model_registry_rejects_unsupported_mode(model_type, spec_mode): ) -def test_dflash_dynamic_verify_uses_fixed_block_token_probabilities(monkeypatch): - block_size = 4 - max_draft_step = 3 - verify_row_count = 5 - flat_draft_token_ids = torch.arange(2 * block_size) - flat_draft_token_probs = torch.arange(2 * block_size, dtype=torch.float32) / 10 - - draft_model = SimpleNamespace( - block_size=block_size, - forward=lambda _: SimpleNamespace(logits=torch.empty(2 * block_size, 1)), - ) - backend = SimpleNamespace( - max_draft_step=max_draft_step, - draft_models=[draft_model], - _gen_argmax_token_ids_and_prob=lambda _: (flat_draft_token_ids, flat_draft_token_probs), - ) - proposer = DFlashProposer(backend=backend, enable_dynmaic_mtp=True) - monkeypatch.setattr( - dflash_module, - "build_parallel_block_draft_input", - lambda **_: (SimpleNamespace(), torch.tensor([10, 11])), - ) - monkeypatch.setattr(dflash_module, "extend_parallel_block_draft_kv_cache", lambda **_: None) - - proposal = proposer.propose_next( - target_model_input=SimpleNamespace(), - target_model_output=SimpleNamespace( - mtp_collector=SimpleNamespace(spec_hidden=torch.empty(verify_row_count, 1)), - ), - target_next_token_ids=torch.arange(verify_row_count), - b_req_mtp_start_loc=torch.tensor([0, 3]), - draft_step=2, - accept_len=torch.tensor([1, 1]), - ) - - assert isinstance(proposal, DFlashSpecProposal) - expected_blocks = flat_draft_token_ids.reshape(2, block_size)[:, :2] - torch.testing.assert_close(proposal.token_ids, expected_blocks) - assert proposal.schedule_scores.shape == (2, 2) - torch.testing.assert_close( - proposal.schedule_scores, - flat_draft_token_probs.reshape(2, block_size)[:, :2], - ) - - -def test_dflash_reuses_decode_input_for_kv_commit(): - forwarded_inputs = [] - draft_model = SimpleNamespace(forward=forwarded_inputs.append) - proposer = DFlashProposer( - backend=SimpleNamespace(draft_models=[draft_model]), - enable_dynmaic_mtp=False, - ) - input_ids = torch.arange(4) - multimodal_params = [{"images": [], "audios": []}] * 4 - position_delta = torch.zeros(4, dtype=torch.int32) - model_input = SimpleNamespace( - batch_size=4, - total_token_num=20, - prefix_total_token_num=None, - input_ids=input_ids, - b_seq_len=torch.tensor([10, 11, 12, 13], dtype=torch.int32), - b_position_delta=position_delta, - multimodal_params=multimodal_params, - is_prefill=False, - ) - target_hidden = torch.empty(4, 16) - - extend_parallel_block_draft_kv_cache( - proposer=proposer, - target_model_input=model_input, - target_hidden=target_hidden, - ) - - assert len(forwarded_inputs) == 1 - assert forwarded_inputs[0] is model_input - assert model_input.input_ids is input_ids - assert model_input.multimodal_params is multimodal_params - assert model_input.total_token_num == 20 - assert model_input.prefix_total_token_num is None - assert not model_input.is_prefill - assert model_input.b_position_delta is position_delta - assert model_input.mtp_draft_input_hiddens is target_hidden - - -def test_dflash_expands_position_delta_with_request_block_rows(monkeypatch): - block_size = 3 - proposer = DFlashProposer( - backend=SimpleNamespace( - draft_models=[SimpleNamespace(block_size=block_size, mask_token_id=99)], - ), - enable_dynmaic_mtp=False, - ) - monkeypatch.setattr( - mtp_utils, - "alloc_mem_indexes", - lambda token_count: torch.arange(token_count, dtype=torch.int32), - ) - model_input = SimpleNamespace( - b_req_idx=torch.tensor([10, 10, 11, 12, 12], dtype=torch.int32), - b_seq_len=torch.tensor([4, 5, 7, 8, 9], dtype=torch.int32), - b_position_delta=torch.tensor([10, 11, 12, 20, 21], dtype=torch.int32), - b_shared_seq_len=torch.tensor([3, 3, 0, 6, 6], dtype=torch.int32), - b_shared_radix_node_id=torch.tensor([1, 1, 2, 3, 3], dtype=torch.int64), - max_kv_seq_len=9, - ) - - draft_input, _ = build_parallel_block_draft_input( - proposer=proposer, - target_model_input=model_input, - target_next_token_ids=torch.arange(5, dtype=torch.int64), - accepted_tail_rows=torch.tensor([1, 4]), - request_count=2, - ) - - assert torch.equal( - draft_input.b_position_delta, - torch.tensor([11, 11, 11, 21, 21, 21], dtype=torch.int32), - ) - assert torch.equal(draft_input.b_seq_len, torch.tensor([6, 7, 8, 10, 11, 12], dtype=torch.int32)) - - def test_hidden_collector_reads_target_layer_ids(monkeypatch): config_reads = [] From 487f448368d66713b8201a548fb136564ce88af6 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Thu, 20 Aug 2026 15:07:20 +0000 Subject: [PATCH 076/103] refactor: derive prefill token count from input ids --- .../basemodel/attention/nsa/flashmla_sparse.py | 2 +- .../basemodel/attention/nsa/fp8_flashmla_sparse.py | 2 +- lightllm/common/basemodel/basemodel.py | 13 +++++++------ lightllm/common/basemodel/prefill_cuda_graph.py | 4 ++-- .../basemodel/test_prefill_cuda_graph_state.py | 6 ++---- 5 files changed, 13 insertions(+), 14 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/flashmla_sparse.py index c3456f4b7a..c25f1d31e7 100644 --- a/lightllm/common/basemodel/attention/nsa/flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/flashmla_sparse.py @@ -51,7 +51,7 @@ def init_state(self): b_q_seq_len=self.infer_state.b_q_seq_len, b_req_idx=self.infer_state.b_req_idx, req_to_token_index=self.infer_state.req_manager.req_to_token_indexs, - q_token_num=self.infer_state.total_token_num - self.infer_state.prefix_total_token_num, + q_token_num=self.infer_state.input_ids.shape[0], ragged_mem_index=self.ragged_mem_index, hold_req_idx=self.infer_state.req_manager.HOLD_REQUEST_ID, ) diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index 539ade769e..4daf7cfc4b 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -46,7 +46,7 @@ def init_state(self): b_q_seq_len=self.infer_state.b_q_seq_len, b_req_idx=self.infer_state.b_req_idx, req_to_token_index=self.infer_state.req_manager.req_to_token_indexs, - q_token_num=self.infer_state.total_token_num - self.infer_state.prefix_total_token_num, + q_token_num=self.infer_state.input_ids.shape[0], ragged_mem_index=self.ragged_mem_index, hold_req_idx=self.infer_state.req_manager.HOLD_REQUEST_ID, ) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index ce35c4158e..f0fff2fb95 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -481,12 +481,13 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s return new_model_input def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle_token_num: int): - if model_input.total_token_num - model_input.prefix_total_token_num == new_handle_token_num: + handle_token_num = model_input.input_ids.shape[0] + if handle_token_num == new_handle_token_num: return model_input - assert model_input.total_token_num - model_input.prefix_total_token_num < new_handle_token_num + assert handle_token_num < new_handle_token_num - padded_token_num = new_handle_token_num - (model_input.total_token_num - model_input.prefix_total_token_num) + padded_token_num = new_handle_token_num - handle_token_num assert padded_token_num > 0 new_model_input = copy.copy(model_input) new_model_input.batch_size = model_input.batch_size + 1 @@ -570,7 +571,7 @@ def _prefill( b_prefill_start_loc=model_input.b_prefill_start_loc, ) - origin_handle_token_num = model_input.total_token_num - model_input.prefix_total_token_num + origin_handle_token_num = model_input.input_ids.shape[0] origin_batch_size = model_input.batch_size if self.args.enable_tpsp_mix_mode: @@ -801,8 +802,8 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod assert model_input1.mem_indexes.is_cuda assert self.args.enable_tpsp_mix_mode - origin_handle_token_num0 = model_input0.total_token_num - model_input0.prefix_total_token_num - origin_handle_token_num1 = model_input1.total_token_num - model_input1.prefix_total_token_num + origin_handle_token_num0 = model_input0.input_ids.shape[0] + origin_handle_token_num1 = model_input1.input_ids.shape[0] infer_handle_token_num0 = triton.cdiv(origin_handle_token_num0, self.tp_world_size_) * self.tp_world_size_ infer_handle_token_num1 = triton.cdiv(origin_handle_token_num1, self.tp_world_size_) * self.tp_world_size_ origin_batch_size0 = model_input0.batch_size diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index 1350e358b3..b1b29aee81 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -77,7 +77,7 @@ def find_closest_graph_handle_token_num(self, handle_token_num: int): def _capture_prefill( self, prefill_func, input_tensors: List[torch.Tensor], infer_state: InferStateInfo ) -> List[torch.Tensor]: - handle_token_num = infer_state.total_token_num - infer_state.prefix_total_token_num + handle_token_num = infer_state.input_ids.shape[0] infer_state.mem_pool = self.mempool infer_state.prefill_cuda_graph_create_graph_obj() infer_state.prefill_cuda_graph_get_current_capture_graph().__enter__() @@ -147,7 +147,7 @@ def capture_prefill( ) def _replay(self, input_tensors: List[torch.Tensor], infer_state: InferStateInfo) -> List[torch.Tensor]: - handle_token_num = infer_state.total_token_num - infer_state.prefix_total_token_num + handle_token_num = infer_state.input_ids.shape[0] graph_infer_state, graph_input_tensors, graph_output_tensors, graph_hidden_collector = self.graph[ handle_token_num ] diff --git a/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py b/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py index 3985e63de8..bed4f3967f 100644 --- a/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py +++ b/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py @@ -10,6 +10,7 @@ class _GraphInferState: def __init__(self, hidden_collector): self.hidden_collector = hidden_collector + self.input_ids = torch.empty(4, dtype=torch.int64) self.copied_from = None self.replayed_with = None @@ -48,8 +49,7 @@ def test_replay_restores_captured_hidden_state_into_request_collector(monkeypatc request_collector = graph_collector.new_instance() graph_infer_state = _GraphInferState(hidden_collector=graph_collector) request_infer_state = SimpleNamespace( - total_token_num=4, - prefix_total_token_num=0, + input_ids=torch.empty(4, dtype=torch.int64), hidden_collector=request_collector, ) graph_output = torch.randn(2, 3) @@ -71,8 +71,6 @@ def test_first_replay_replaces_capture_collector_with_runtime_instance(monkeypat graph_collector = _create_layer_collector(monkeypatch) graph_collector.add(layer_index=0, hidden=torch.randn(2, 3)) graph_infer_state = _GraphInferState(hidden_collector=graph_collector) - graph_infer_state.total_token_num = 4 - graph_infer_state.prefix_total_token_num = 0 graph_output = torch.randn(2, 3) prefill_graph = PrefillCudaGraph.__new__(PrefillCudaGraph) prefill_graph.graph = {4: (graph_infer_state, [], [graph_output], graph_collector)} From 144cd30d51e8eca8034990e3c9ac9707b44bb3e9 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 21 Aug 2026 02:14:44 +0000 Subject: [PATCH 077/103] refactor dp padding into base model --- lightllm/common/basemodel/basemodel.py | 96 ++++++------- lightllm/common/basemodel/batch_objs.py | 7 +- .../mode_backend/dp_backend/impl.py | 22 +-- .../mode_backend/generic_pre_process.py | 23 ++-- .../common/basemodel/test_model_input.py | 86 ++++++++++++ .../common/basemodel/test_model_output.py | 126 +++++++++++++++++- .../mode_backend/test_generic_pre_process.py | 48 ++++++- 7 files changed, 335 insertions(+), 73 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index f0fff2fb95..18ea77cfd2 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -561,7 +561,7 @@ def _prefill( self, model_input: ModelInput, ): - if self.args.enable_prefill_decode_mixed: + if self.args.enable_prefill_decode_mixed and model_input.input_ids.shape[0] > 0: gather_token_prefill_decode_mixed( input_ids=model_input.input_ids, req_to_next_token_ids=self.req_manager.req_sampling_params_manager.req_to_next_token_ids, @@ -574,10 +574,11 @@ def _prefill( origin_handle_token_num = model_input.input_ids.shape[0] origin_batch_size = model_input.batch_size + # 即使当前 DP rank 没有实际 token,也需要补出一个 dummy token 执行模型; + # TPSP 模式再将这个实际推理数量向上对齐到 TP world size 的整数倍。 + infer_handle_token_num = max(1, origin_handle_token_num) if self.args.enable_tpsp_mix_mode: - infer_handle_token_num = triton.cdiv(origin_handle_token_num, self.tp_world_size_) * self.tp_world_size_ - else: - infer_handle_token_num = origin_handle_token_num + infer_handle_token_num = triton.cdiv(infer_handle_token_num, self.tp_world_size_) * self.tp_world_size_ if self.prefill_graph is not None and self.prefill_graph.can_run(handle_token_num=infer_handle_token_num): infer_handle_token_num = self.prefill_graph.find_closest_graph_handle_token_num( @@ -617,62 +618,64 @@ def _decode( self, model_input: ModelInput, ) -> ModelOutput: - # for overlap mode if model_input.input_ids is None: - model_input.input_ids = gather_token( - self.req_manager.req_sampling_params_manager.req_to_next_token_ids, - model_input.b_req_idx, - model_input.b_mtp_index, - ) + if model_input.batch_size > 0: + model_input.input_ids = gather_token( + self.req_manager.req_sampling_params_manager.req_to_next_token_ids, + model_input.b_req_idx, + model_input.b_mtp_index, + ) + else: + # 空 DP rank 不启动 gather kernel,但仍将 input_ids 规范化为 + # CUDA 空 tensor;后续内部 padding 会为 dummy request 填入 token id 1。 + model_input.input_ids = torch.empty( + (0,), + dtype=torch.int64, + device=model_input.b_req_idx.device, + ) origin_batch_size = model_input.batch_size + # 空 DP rank 先补出一个 dummy request;TPSP 模式下继续将 batch size + # 向上对齐到 TP world size 的整数倍,保证后续切分得到合法 shape。 + infer_batch_size = max(1, origin_batch_size) if self.args.enable_tpsp_mix_mode: - infer_batch_size = triton.cdiv(model_input.batch_size, self.tp_world_size_) * self.tp_world_size_ - else: - infer_batch_size = model_input.batch_size - - if self.graph is not None and self.graph.can_run( - batch_size=infer_batch_size, max_len_in_batch=model_input.max_kv_seq_len - ): + infer_batch_size = triton.cdiv(infer_batch_size, self.tp_world_size_) * self.tp_world_size_ + + # CUDA Graph 可能继续向上对齐 batch size,并因此加入 seq_len=2 的 + # dummy request。先用最终可能出现的 KV 长度判断 graph,再统一 padding 一次。 + infer_max_kv_seq_len = max(2, model_input.max_kv_seq_len) + use_cuda_graph = self.graph is not None and self.graph.can_run( + batch_size=infer_batch_size, + max_len_in_batch=infer_max_kv_seq_len, + ) + need_capture = False + if use_cuda_graph: infer_batch_size = self.graph.find_closest_graph_batch_size(batch_size=infer_batch_size) - model_input = self._create_padded_decode_model_input( - model_input=model_input, new_batch_size=infer_batch_size - ) - infer_state = self._create_inferstate(model_input) need_capture = self.graph.need_capture(infer_batch_size) - infer_state.is_cuda_graph = need_capture - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state.b_req_idx, - infer_state.b_seq_len, - infer_state.mem_index, - ) - infer_state.init_some_extra_state(self) - infer_state.init_att_state() + model_input = self._create_padded_decode_model_input(model_input=model_input, new_batch_size=infer_batch_size) + infer_state = self._create_inferstate(model_input) + # attention backend 会根据该标记准备 CUDA Graph capture 专用状态, + # 因此必须在 init_att_state 之前完成赋值。 + infer_state.is_cuda_graph = need_capture + copy_kv_index_to_req( + self.req_manager.req_to_token_indexs, + infer_state.b_req_idx, + infer_state.b_seq_len, + infer_state.mem_index, + ) + infer_state.init_some_extra_state(self) + infer_state.init_att_state() + + if use_cuda_graph: if need_capture: model_output: ModelOutput = self.graph.capture_decode(self._token_forward, infer_state) else: model_output: ModelOutput = self.graph.replay(infer_state) - - model_output = self._create_unpad_decode_model_output(model_output, origin_batch_size=origin_batch_size) else: - model_input = self._create_padded_decode_model_input( - model_input=model_input, new_batch_size=infer_batch_size - ) - infer_state = self._create_inferstate(model_input) - copy_kv_index_to_req( - self.req_manager.req_to_token_indexs, - infer_state.b_req_idx, - infer_state.b_seq_len, - infer_state.mem_index, - ) - infer_state.init_some_extra_state(self) - infer_state.init_att_state() model_output = self._token_forward(infer_state) - model_output = self._create_unpad_decode_model_output(model_output, origin_batch_size=origin_batch_size) - return model_output + return self._create_unpad_decode_model_output(model_output, origin_batch_size=origin_batch_size) @final def _context_forward(self, infer_state: InferStateInfo): @@ -1119,6 +1122,7 @@ def _check_max_len_infer(self): is_prefill=True, b_ready_cache_len=b_ready_cache_len, b_prefill_start_loc=b_prefill_start_loc, + b_prefill_has_output_cpu=[False], multimodal_params=[{"images": [], "audios": []}], ) model_output = self.forward( diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 93912dfe08..3da53b04c0 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -46,7 +46,8 @@ class ModelInput: mem_indexes_cpu: torch.Tensor = None # prefill 阶段使用的参数,但是不是推理过程使用的参数,是推理外部进行资源管理 # 的一些变量 - b_prefill_has_output_cpu: List[bool] = None # 标记进行prefill的请求是否具有输出 + # 标记 prefill 请求是否会在本轮产生输出。Prefill 必填(空 batch 使用空 list),decode 不使用。 + b_prefill_has_output_cpu: List[bool] = None # 专有变量,用于一些特殊的模型,特殊的模式下, 传递一些特殊 # 的输入变量。只在特殊的模型模式下才会具体使用和生效。 @@ -113,8 +114,8 @@ def check_input(self): assert self.b_is_decode_req.shape == self.b_req_idx.shape assert self.b_is_decode_req.dtype == torch.bool assert self.b_position_delta is None, "prefill must not provide b_position_delta" - if self.b_prefill_has_output_cpu is not None: - assert len(self.b_prefill_has_output_cpu) == self.batch_size + assert self.b_prefill_has_output_cpu is not None, "prefill must provide b_prefill_has_output_cpu" + assert len(self.b_prefill_has_output_cpu) == self.batch_size else: assert self.max_q_seq_len == 1 assert self.b_position_delta is not None diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 49be5fcced..9b011aad6a 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -6,6 +6,8 @@ from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample from lightllm.server.router.model_infer.mode_backend.pre import ( + prepare_prefill_inputs, + prepare_decode_inputs, padded_prepare_prefill_inputs, padded_prepare_decode_inputs, padded_overlap_prepare_prefill_inputs, @@ -198,7 +200,7 @@ def prefill_normal( event_pack: OverlapEventPack, prefill_reqs: List[InferReq], ): - model_input, run_reqs, _ = padded_prepare_prefill_inputs(prefill_reqs) + model_input, run_reqs = prepare_prefill_inputs(prefill_reqs, is_chuncked_mode=not self.disable_chunked_prefill) run_reqs_num = len(run_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output = self.model.forward(model_input) @@ -210,16 +212,16 @@ def prefill_normal( next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=model_output.logits[:run_reqs_num], - b_req_idx=model_input.b_req_idx[:run_reqs_num], - b_mtp_index=model_input.b_mtp_index[:run_reqs_num], + logits=model_output.logits, + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, run_reqs=run_reqs, is_prefill=True, - b_prefill_has_output_cpu=model_input.b_prefill_has_output_cpu[:run_reqs_num], + b_prefill_has_output_cpu=model_input.b_prefill_has_output_cpu, mask_func=None, ) g_infer_context.copy_linear_att_state_to_cache_buffer( - b_req_idx=model_input.b_req_idx[:run_reqs_num], + b_req_idx=model_input.b_req_idx, reqs=run_reqs, ) sync_event = torch.cuda.Event() @@ -251,7 +253,7 @@ def prefill_normal( return def decode_normal(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): - model_input, run_reqs, padded_req_num = padded_prepare_decode_inputs(req_objs=decode_reqs) + model_input, run_reqs = prepare_decode_inputs(req_objs=decode_reqs) model_input: ModelInput = model_input run_reqs_num = len(run_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): @@ -263,9 +265,9 @@ def decode_normal(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=model_output.logits[:run_reqs_num], - b_req_idx=model_input.b_req_idx[:run_reqs_num], - b_mtp_index=model_input.b_mtp_index[:run_reqs_num], + logits=model_output.logits, + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, run_reqs=run_reqs, is_prefill=False, mask_func=None, diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 6d9920b585..40b1747714 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -51,11 +51,13 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> else: b_is_decode_req.append(False) - max_kv_seq_len = max(b_seq_len) - max_cache_len = max(b_ready_cache_len) - max_q_seq_len = max(b_q_seq_len) + # DP 模式下某个 rank 可能没有本地请求。这里保留真实的 0 shape, + # 推理所需的 dummy request 统一由 BaseModel 在执行前补齐。 + max_kv_seq_len = max(b_seq_len, default=0) + max_cache_len = max(b_ready_cache_len, default=0) + max_q_seq_len = max(b_q_seq_len, default=0) - input_ids = np.concatenate(input_ids, dtype=np.int64) + input_ids = np.concatenate(input_ids, dtype=np.int64) if input_ids else np.empty((0,), dtype=np.int64) input_ids = torch.tensor(input_ids, dtype=torch.int64, device="cpu") b_req_idx = torch.tensor(b_req_idx, dtype=torch.int32, device="cpu") b_seq_len = torch.tensor(b_seq_len, dtype=torch.int32, device="cpu") @@ -122,8 +124,10 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In multimodal_params.append(req.multimodal_params) b_q_seq_len.append(1) - max_kv_seq_len = max(b_seq_len) - max_q_seq_len = max(b_q_seq_len) + # 空 DP rank 同样构建完整的 decode ModelInput;BaseModel 会在 token + # gather 和 attention 初始化之前补入内部 dummy request。 + max_kv_seq_len = max(b_seq_len, default=0) + max_q_seq_len = max(b_q_seq_len, default=1) b_req_idx = torch.tensor(b_req_idx, dtype=torch.int32, device="cpu") b_seq_len = torch.tensor(b_seq_len, dtype=torch.int32, device="cpu") @@ -134,12 +138,7 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In [req.get_radix_cache_shared_len() for req in run_reqs], dtype=torch.int32, device="cpu" ) b_shared_radix_node_id = torch.tensor( - [ - -1 - if req.shared_kv_node is None - else req.shared_kv_node.time_id % INT64_MAX - for req in run_reqs - ], + [-1 if req.shared_kv_node is None else req.shared_kv_node.time_id % INT64_MAX for req in run_reqs], dtype=torch.int64, device="cpu", ) diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index 58116c9ff0..f6d7901ac0 100644 --- a/unit_tests/common/basemodel/test_model_input.py +++ b/unit_tests/common/basemodel/test_model_input.py @@ -83,6 +83,14 @@ def test_prefill_requires_prefill_metadata(): model_input.check_input() +def test_prefill_requires_prefill_output_markers(): + model_input = _create_model_input(is_prefill=True) + model_input.b_prefill_has_output_cpu = None + + with pytest.raises(AssertionError, match="prefill must provide b_prefill_has_output_cpu"): + model_input.check_input() + + def test_prefill_requires_decode_request_markers(): model_input = _create_model_input(is_prefill=True) model_input.b_is_decode_req = None @@ -132,3 +140,81 @@ def test_padded_prefill_adds_non_decode_request_marker(): assert padded_input.b_is_decode_req.dtype == torch.bool assert padded_input.b_is_decode_req.tolist() == [True, False] + + +def test_padded_prefill_builds_internal_request_for_empty_input(): + model_input = ModelInput( + batch_size=0, + total_token_num=0, + max_q_seq_len=0, + max_kv_seq_len=0, + max_cache_len=0, + prefix_total_token_num=0, + input_ids=torch.empty((0,), dtype=torch.int64), + b_req_idx=torch.empty((0,), dtype=torch.int32), + b_mtp_index=torch.empty((0,), dtype=torch.int32), + b_seq_len=torch.empty((0,), dtype=torch.int32), + b_is_decode_req=torch.empty((0,), dtype=torch.bool), + b_ready_cache_len=torch.empty((0,), dtype=torch.int32), + b_prefill_start_loc=torch.empty((0,), dtype=torch.int32), + mem_indexes=torch.empty((0,), dtype=torch.int32), + is_prefill=True, + b_prefill_has_output_cpu=[], + multimodal_params=[], + ) + model = SimpleNamespace( + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + req_manager=SimpleNamespace(HOLD_REQUEST_ID=88), + ) + + padded_input = TpPartBaseModel._create_padded_prefill_model_input( + model, + model_input=model_input, + new_handle_token_num=1, + ) + + assert model_input.batch_size == 0 + assert padded_input.batch_size == 1 + assert padded_input.input_ids.tolist() == [1] + assert padded_input.mem_indexes.tolist() == [99] + assert padded_input.b_req_idx.tolist() == [88] + assert padded_input.b_seq_len.tolist() == [1] + assert padded_input.b_prefill_has_output_cpu == [False] + + +def test_padded_decode_builds_internal_request_from_empty_token_tensor(): + model_input = ModelInput( + batch_size=0, + total_token_num=0, + max_q_seq_len=1, + max_kv_seq_len=0, + input_ids=torch.empty((0,), dtype=torch.int64), + b_req_idx=torch.empty((0,), dtype=torch.int32), + b_mtp_index=torch.empty((0,), dtype=torch.int32), + b_seq_len=torch.empty((0,), dtype=torch.int32), + b_position_delta=torch.empty((0,), dtype=torch.int32), + b_shared_seq_len=torch.empty((0,), dtype=torch.int32), + b_shared_radix_node_id=torch.empty((0,), dtype=torch.int64), + mem_indexes=torch.empty((0,), dtype=torch.int32), + is_prefill=False, + multimodal_params=[], + ) + model = SimpleNamespace( + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + req_manager=SimpleNamespace(HOLD_REQUEST_ID=88), + ) + + padded_input = TpPartBaseModel._create_padded_decode_model_input( + model, + model_input=model_input, + new_batch_size=1, + ) + + assert model_input.batch_size == 0 + assert padded_input.batch_size == 1 + assert padded_input.input_ids.tolist() == [1] + assert padded_input.mem_indexes.tolist() == [99] + assert padded_input.b_req_idx.tolist() == [88] + assert padded_input.b_seq_len.tolist() == [2] + assert padded_input.b_shared_seq_len.tolist() == [0] + assert padded_input.b_shared_radix_node_id.tolist() == [-1] diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index 18333d763c..6f9477e294 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -1,7 +1,10 @@ +from types import SimpleNamespace + import torch +from lightllm.common.basemodel import basemodel from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput +from lightllm.common.basemodel.batch_objs import ModelInput, ModelMtpOutputCollector, ModelOutput def test_decode_unpad_slices_spec_output_with_logits(): @@ -38,3 +41,124 @@ def test_prefill_unpad_uses_token_rows_for_spec_hidden(): assert unpadded.logits.shape == (3, 4) assert unpadded.mtp_collector.spec_hidden.shape == (6, 3) assert unpadded.prompt_logics.shape == (6, 4) + + +def test_decode_unpad_restores_empty_output(): + model = TpPartBaseModel.__new__(TpPartBaseModel) + output = ModelOutput( + logits=torch.arange(4).view(1, 4), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.arange(3).view(1, 3)), + ) + + unpadded = model._create_unpad_decode_model_output(output, origin_batch_size=0) + + assert unpadded.logits.shape == (0, 4) + assert unpadded.mtp_collector.spec_hidden.shape == (0, 3) + + +def test_prefill_unpad_restores_empty_output(): + model = TpPartBaseModel.__new__(TpPartBaseModel) + output = ModelOutput( + logits=torch.arange(4).view(1, 4), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.arange(6).view(2, 3)), + prompt_logics=torch.arange(8).view(2, 4), + ) + + unpadded = model._create_unpad_prefill_model_output( + output, + origin_handle_token_num=0, + origin_batch_size=0, + ) + + assert unpadded.logits.shape == (0, 4) + assert unpadded.mtp_collector.spec_hidden.shape == (0, 3) + assert unpadded.prompt_logics.shape == (0, 4) + + +def _create_empty_decode_input(): + return ModelInput( + batch_size=0, + total_token_num=0, + max_q_seq_len=1, + max_kv_seq_len=0, + input_ids=torch.empty((0,), dtype=torch.int64), + b_req_idx=torch.empty((0,), dtype=torch.int32), + b_mtp_index=torch.empty((0,), dtype=torch.int32), + b_seq_len=torch.empty((0,), dtype=torch.int32), + b_position_delta=torch.empty((0,), dtype=torch.int32), + b_shared_seq_len=torch.empty((0,), dtype=torch.int32), + b_shared_radix_node_id=torch.empty((0,), dtype=torch.int64), + mem_indexes=torch.empty((0,), dtype=torch.int32), + is_prefill=False, + multimodal_params=[], + ) + + +@torch.no_grad() +def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): + monkeypatch.setattr(basemodel, "copy_kv_index_to_req", lambda *args: None) + + execution_configs = ( + # eager 普通模式:空 batch 只补一个 dummy request。 + (None, False, 1, False, 1), + # eager TPSP 模式:dummy request 仍需对齐到 TP world size。 + (None, True, 2, False, 2), + # CUDA Graph replay:使用最终选中的 graph batch size。 + (4, False, 1, False, 4), + # CUDA Graph capture:必须在初始化 attention state 前设置 graph 标记。 + (4, False, 1, True, 4), + ) + for ( + graph_batch_size, + enable_tpsp_mix_mode, + tp_world_size, + need_capture, + expected_batch_size, + ) in execution_configs: + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.args = SimpleNamespace(enable_tpsp_mix_mode=enable_tpsp_mix_mode) + model.tp_world_size_ = tp_world_size + model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=99) + model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) + + pad_batch_sizes = [] + + def pad_once(model_input, new_batch_size): + pad_batch_sizes.append(new_batch_size) + return TpPartBaseModel._create_padded_decode_model_input(model, model_input, new_batch_size) + + model._create_padded_decode_model_input = pad_once + + graph_flags_at_att_init = [] + + def create_infer_state(model_input): + infer_state = SimpleNamespace( + b_req_idx=model_input.b_req_idx, + b_seq_len=model_input.b_seq_len, + mem_index=model_input.mem_indexes, + init_some_extra_state=lambda _: None, + is_cuda_graph=False, + ) + infer_state.init_att_state = lambda: graph_flags_at_att_init.append(infer_state.is_cuda_graph) + return infer_state + + model._create_inferstate = create_infer_state + model._token_forward = lambda infer_state: ModelOutput(logits=torch.ones((infer_state.b_req_idx.shape[0], 4))) + + graph = None + if graph_batch_size is not None: + graph = SimpleNamespace() + graph.can_run = lambda **kwargs: True + graph.find_closest_graph_batch_size = lambda batch_size: graph_batch_size + graph.need_capture = lambda batch_size: need_capture + graph.capture_decode = lambda decode_func, infer_state: ModelOutput( + logits=torch.ones((infer_state.b_req_idx.shape[0], 4)) + ) + graph.replay = lambda infer_state: ModelOutput(logits=torch.ones((infer_state.b_req_idx.shape[0], 4))) + model.graph = graph + + output = model._decode(_create_empty_decode_input()) + + assert pad_batch_sizes == [expected_batch_size] + assert graph_flags_at_att_init == [need_capture] + assert output.logits.shape == (0, 4) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index 6ed1324ec6..ad9d55d2ec 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -1,7 +1,53 @@ from types import SimpleNamespace import torch -from lightllm.server.router.model_infer.mode_backend import generic_padded_pre_process +from lightllm.server.router.model_infer.mode_backend import generic_padded_pre_process, generic_pre_process + + +def _patch_empty_input_context(monkeypatch): + mem_manager = SimpleNamespace( + HOLD_TOKEN_MEMINDEX=-1, + alloc=lambda size: torch.empty((size,), dtype=torch.int32), + ) + infer_context = SimpleNamespace( + req_manager=SimpleNamespace(HOLD_REQUEST_ID=-1, mem_manager=mem_manager), + radix_cache=None, + ) + monkeypatch.setattr(generic_pre_process, "g_infer_context", infer_context) + return infer_context + + +def test_prepare_prefill_inputs_allows_empty_batch(monkeypatch): + _patch_empty_input_context(monkeypatch) + + model_input, run_reqs = generic_pre_process.prepare_prefill_inputs([], is_chuncked_mode=True) + + assert run_reqs == [] + assert model_input.batch_size == 0 + assert model_input.input_ids.shape == (0,) + assert model_input.mem_indexes_cpu.shape == (0,) + assert model_input.b_req_idx.shape == (0,) + assert model_input.b_prefill_start_loc.shape == (0,) + assert model_input.b_prefill_has_output_cpu == [] + assert model_input.max_q_seq_len == 0 + assert model_input.max_kv_seq_len == 0 + + +def test_prepare_decode_inputs_allows_empty_batch(monkeypatch): + _patch_empty_input_context(monkeypatch) + + model_input, run_reqs = generic_pre_process.prepare_decode_inputs([]) + + assert run_reqs == [] + assert model_input.batch_size == 0 + assert model_input.input_ids is None + assert model_input.mem_indexes_cpu.shape == (0,) + assert model_input.b_req_idx.shape == (0,) + assert model_input.b_position_delta.shape == (0,) + assert model_input.b_shared_seq_len.shape == (0,) + assert model_input.b_shared_radix_node_id.shape == (0,) + assert model_input.max_q_seq_len == 1 + assert model_input.max_kv_seq_len == 0 def test_padded_decode_builds_raw_shared_radix_metadata(monkeypatch): From a0735188f859ea2fbd7e1209c6bbd79fbc2a70d9 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 21 Aug 2026 02:26:48 +0000 Subject: [PATCH 078/103] remove redundant prefix token count --- .../basemodel/attention/nsa/fp8_flashmla_sparse.py | 2 +- lightllm/common/basemodel/basemodel.py | 5 ----- lightllm/common/basemodel/batch_objs.py | 2 -- lightllm/common/basemodel/infer_struct.py | 3 --- lightllm/common/basemodel/prefill_cuda_graph.py | 2 -- .../mode_backend/generic_padded_pre_process.py | 11 +---------- .../model_infer/mode_backend/generic_pre_process.py | 3 --- test/benchmark/static_inference/static_benchmark.py | 2 -- unit_tests/common/basemodel/test_model_input.py | 3 --- .../model_infer/mtp_speculative/test_eagle_overlap.py | 1 - 10 files changed, 2 insertions(+), 32 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index 4daf7cfc4b..c58f88244e 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -79,7 +79,7 @@ def _nsa_prefill_att( topk_mem_indices = nsa_dict["topk_mem_indices"] prefill_cache_kv = nsa_dict["prefill_cache_kv"] - if self.infer_state.prefix_total_token_num > 0: + if self.infer_state.max_cache_len > 0: # 当前推理生成的token kv部分从 prefill_cache_kv 中获取,历史 # 部分kv 从 packed_kv 中获取, 并进行反量化,这样可以避免 prefill_cache_kv # 部分的数据进行重复的反量化操作,提升整体的性能。 diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 18ea77cfd2..f4a59ef278 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -388,7 +388,6 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) infer_state.max_q_seq_len = model_input.max_q_seq_len infer_state.max_kv_seq_len = model_input.max_kv_seq_len infer_state.max_cache_len = model_input.max_cache_len - infer_state.prefix_total_token_num = model_input.prefix_total_token_num assert model_input.b_req_idx.shape[0] == model_input.b_seq_len.shape[0] infer_state.b_req_idx = model_input.b_req_idx infer_state.b_seq_len = model_input.b_seq_len @@ -513,7 +512,6 @@ def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle new_model_input.b_prefill_start_loc = b_q_seq_len.cumsum(dim=0, dtype=torch.int32) - b_q_seq_len # 构建新的list, 使用 append 可能会让外面使用的数组引用发生变化,导致错误。 new_model_input.b_prefill_has_output_cpu = [e for e in new_model_input.b_prefill_has_output_cpu] + [False] - new_model_input.prefix_total_token_num = model_input.prefix_total_token_num new_model_input.multimodal_params = [e for e in new_model_input.multimodal_params] + [ {"images": [], "audios": []} @@ -1112,7 +1110,6 @@ def _check_max_len_infer(self): max_q_seq_len=self.batch_max_tokens, max_kv_seq_len=self.batch_max_tokens, max_cache_len=0, - prefix_total_token_num=0, input_ids=dummy_input_ids, mem_indexes=mem_indexes, b_req_idx=b_req_idx, @@ -1192,7 +1189,6 @@ def _autotune_warmup(self): max_q_seq_len=input_len, max_kv_seq_len=input_len, max_cache_len=0, - prefix_total_token_num=0, input_ids=dummy_input_ids, mem_indexes=mem_indexes, b_req_idx=b_req_idx, @@ -1259,7 +1255,6 @@ def _init_padded_req(self): max_q_seq_len=prefill_input_len, max_kv_seq_len=prefill_input_len, max_cache_len=0, - prefix_total_token_num=0, input_ids=dummy_input_ids, mem_indexes=mem_indexes, b_req_idx=b_req_idx, diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 3da53b04c0..ae645d4b7b 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -16,7 +16,6 @@ class ModelInput: max_q_seq_len: int max_kv_seq_len: int max_cache_len: int = None - prefix_total_token_num: int = None input_ids: torch.Tensor = None b_req_idx: torch.Tensor = None b_mtp_index: torch.Tensor = None @@ -105,7 +104,6 @@ def check_input(self): if self.is_prefill: assert self.input_ids is not None assert self.max_cache_len is not None - assert self.prefix_total_token_num is not None assert self.b_ready_cache_len is not None assert self.b_prefill_start_loc is not None assert self.b_is_decode_req is not None diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index 6434102dbe..91c6e99699 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -44,9 +44,6 @@ def __init__(self): self.b_seq_len: torch.Tensor = None # max_cache_len 用于 prefill 阶段标识请求中最大 cache的kv 的长度 self.max_cache_len: int = None - # prefix_total_token_num 用于 prefill 阶段标识当前请求中所有已经ready的kv的长度 - # 的sum值, 其值等于 sum(b_ready_cache_len) - self.prefix_total_token_num: int = None self.is_prefill: bool = None self.mem_manager: MemoryManager = None diff --git a/lightllm/common/basemodel/prefill_cuda_graph.py b/lightllm/common/basemodel/prefill_cuda_graph.py index b1b29aee81..bf6039a48f 100644 --- a/lightllm/common/basemodel/prefill_cuda_graph.py +++ b/lightllm/common/basemodel/prefill_cuda_graph.py @@ -219,7 +219,6 @@ def warmup(self, model): b_prefill_start_loc=b_prefill_start_loc, is_prefill=True, b_prefill_has_output_cpu=[False], - prefix_total_token_num=0, multimodal_params=[{"images": [], "audios": []}], **model._gen_special_model_input(token_num=total_token_num), ) @@ -281,7 +280,6 @@ def warmup_overlap(self, model): b_prefill_start_loc=b_prefill_start_loc, is_prefill=True, b_prefill_has_output_cpu=[False], - prefix_total_token_num=0, multimodal_params=[{"images": [], "audios": []}], **model._gen_special_model_input(token_num=total_token_num), ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py index d829163fb1..0e4a562348 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py @@ -29,7 +29,6 @@ def padded_prepare_prefill_inputs( run_reqs = [] total_token_num = 0 - prefix_total_token_num = 0 padded_req_num = dest_batch_size - len(req_objs) input_ids = [] b_req_idx = [] @@ -57,7 +56,6 @@ def padded_prepare_prefill_inputs( b_q_seq_len.append(input_token_len) input_ids.append(input_id) total_token_num += seq_len - prefix_total_token_num += req.cur_kv_len b_ready_cache_len.append(req.cur_kv_len) b_mtp_index.append(0) @@ -79,7 +77,6 @@ def padded_prepare_prefill_inputs( b_prefill_has_output.append(False) b_ready_cache_len.append(0) total_token_num += 1 - prefix_total_token_num += 0 batch_multimodal_params.append({"images": [], "audios": []}) b_is_decode_req.append(False) @@ -116,7 +113,6 @@ def padded_prepare_prefill_inputs( max_q_seq_len=max_q_seq_len, max_kv_seq_len=max_kv_seq_len, max_cache_len=max_cache_len, - prefix_total_token_num=prefix_total_token_num, input_ids=input_ids, mem_indexes_cpu=mem_indexes, b_req_idx=b_req_idx, @@ -207,12 +203,7 @@ def padded_prepare_decode_inputs( [req.get_radix_cache_shared_len() for req in run_reqs], dtype=torch.int32, device="cpu" ) b_shared_radix_node_id = torch.tensor( - [ - -1 - if req.shared_kv_node is None - else req.shared_kv_node.time_id % INT64_MAX - for req in run_reqs - ], + [-1 if req.shared_kv_node is None else req.shared_kv_node.time_id % INT64_MAX for req in run_reqs], dtype=torch.int64, device="cpu", ) diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 40b1747714..313aea64f1 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -10,7 +10,6 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> Tuple[ModelInput, List[InferReq]]: run_reqs = [] total_token_num = 0 - prefix_total_token_num = 0 input_ids = [] b_req_idx = [] b_seq_len = [] @@ -42,7 +41,6 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> b_q_seq_len.append(input_token_len) input_ids.append(input_id) total_token_num += seq_len - prefix_total_token_num += req.cur_kv_len b_ready_cache_len.append(req.cur_kv_len) b_mtp_index.append(0) if hasattr(req, "is_decode_req_mixed_in_prefill"): @@ -88,7 +86,6 @@ def prepare_prefill_inputs(req_objs: List[InferReq], is_chuncked_mode: bool) -> b_prefill_start_loc=b_prefill_start_loc, is_prefill=True, b_prefill_has_output_cpu=b_prefill_has_output, - prefix_total_token_num=prefix_total_token_num, multimodal_params=batch_multimodal_params, ) diff --git a/test/benchmark/static_inference/static_benchmark.py b/test/benchmark/static_inference/static_benchmark.py index b3faf99130..b3f19f360a 100644 --- a/test/benchmark/static_inference/static_benchmark.py +++ b/test/benchmark/static_inference/static_benchmark.py @@ -539,7 +539,6 @@ def _make_prefill_input(self, token_chunk: np.ndarray, req_idx: torch.Tensor, re max_q_seq_len=q_len, max_kv_seq_len=seq_len_value, max_cache_len=ready_cache_len, - prefix_total_token_num=ready_cache_len * batch_size, input_ids=input_ids, b_req_idx=req_idx, b_mtp_index=cpu_i32_zeros(batch_size), @@ -626,7 +625,6 @@ def _slice_prefill_input( max_q_seq_len=int(b_q_seq_len.max().item()), max_kv_seq_len=int(b_seq_len.max().item()), max_cache_len=int(b_ready_cache_len.max().item()), - prefix_total_token_num=int(b_ready_cache_len.sum().item()), input_ids=model_input.input_ids[token_start:token_end].contiguous(), b_req_idx=model_input.b_req_idx[batch_start:batch_end].clone(), b_mtp_index=model_input.b_mtp_index[batch_start:batch_end].clone(), diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index f6d7901ac0..9853d78cfa 100644 --- a/unit_tests/common/basemodel/test_model_input.py +++ b/unit_tests/common/basemodel/test_model_input.py @@ -23,7 +23,6 @@ def _create_model_input(*, is_prefill=False): ) if is_prefill: kwargs["max_cache_len"] = 0 - kwargs["prefix_total_token_num"] = 0 kwargs["input_ids"] = torch.ones(batch_size, dtype=torch.int64) kwargs["b_ready_cache_len"] = torch.zeros(batch_size, dtype=torch.int32) kwargs["b_prefill_start_loc"] = torch.arange(batch_size, dtype=torch.int32) @@ -114,7 +113,6 @@ def test_padded_prefill_adds_non_decode_request_marker(): max_q_seq_len=2, max_kv_seq_len=2, max_cache_len=0, - prefix_total_token_num=0, input_ids=torch.ones(2, dtype=torch.int64), b_req_idx=torch.zeros(1, dtype=torch.int32), b_mtp_index=torch.zeros(1, dtype=torch.int32), @@ -149,7 +147,6 @@ def test_padded_prefill_builds_internal_request_for_empty_input(): max_q_seq_len=0, max_kv_seq_len=0, max_cache_len=0, - prefix_total_token_num=0, input_ids=torch.empty((0,), dtype=torch.int64), b_req_idx=torch.empty((0,), dtype=torch.int32), b_mtp_index=torch.empty((0,), dtype=torch.int32), diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 77def5b36a..99c563a3f9 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -48,7 +48,6 @@ def _target_input(batch_size): return SimpleNamespace( batch_size=batch_size, total_token_num=batch_size, - prefix_total_token_num=None, input_ids=torch.arange(batch_size, dtype=torch.int64), b_seq_len=torch.arange(batch_size, dtype=torch.int32) + 4, b_req_idx=torch.arange(batch_size, dtype=torch.int32), From 4932e3300b9b7784143ed217cafb410bcc79e06f Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 21 Aug 2026 08:13:04 +0000 Subject: [PATCH 079/103] refactor: simplify overlap input preparation --- lightllm/common/basemodel/basemodel.py | 109 +++---- .../triton_kernel/gather_token_id.py | 6 +- .../triton_kernel/gen_mtp_prefill_params.py | 4 + .../mode_backend/dp_backend/impl.py | 88 +++--- .../generic_padded_pre_process.py | 270 ------------------ .../mode_backend/generic_pre_process.py | 48 ++++ .../router/model_infer/mode_backend/pre.py | 12 +- .../mtp_speculative/dp_overlap_engine.py | 12 +- .../dp_overlap_proposers/eagle_utils.py | 8 +- .../dp_overlap_proposers/vanilla_no_att.py | 2 +- .../dp_overlap_proposers/vanilla_with_att.py | 4 +- .../static_inference/static_benchmark.py | 100 +++++-- .../common/basemodel/test_overlap_utils.py | 213 ++++++++++++++ .../test_gen_mtp_prefill_params.py | 16 ++ .../mode_backend/test_generic_pre_process.py | 136 ++++++--- .../mtp_speculative/test_eagle_overlap.py | 4 +- .../mtp_speculative/test_vanilla_overlap.py | 2 +- .../mtp_speculative/test_vanilla_prefill.py | 2 +- 18 files changed, 576 insertions(+), 460 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py create mode 100644 unit_tests/common/basemodel/test_overlap_utils.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index f4a59ef278..dbeb15ffde 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -619,9 +619,9 @@ def _decode( if model_input.input_ids is None: if model_input.batch_size > 0: model_input.input_ids = gather_token( - self.req_manager.req_sampling_params_manager.req_to_next_token_ids, - model_input.b_req_idx, - model_input.b_mtp_index, + req_to_next_token_ids=(self.req_manager.req_sampling_params_manager.req_to_next_token_ids), + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, ) else: # 空 DP rank 不启动 gather kernel,但仍将 input_ids 规范化为 @@ -776,37 +776,36 @@ def _token_forward(self, infer_state: InferStateInfo): @torch.no_grad() def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: ModelInput): - model_input0.to_cuda() - model_input1.to_cuda() - - if self.args.enable_prefill_decode_mixed: - gather_token_prefill_decode_mixed( - input_ids=model_input0.input_ids, - req_to_next_token_ids=self.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_idx=model_input0.b_req_idx, - b_mtp_index=model_input0.b_mtp_index, - b_is_decode_req=model_input0.b_is_decode_req, - b_prefill_start_loc=model_input0.b_prefill_start_loc, - ) - - if self.args.enable_prefill_decode_mixed: - gather_token_prefill_decode_mixed( - input_ids=model_input1.input_ids, - req_to_next_token_ids=self.req_manager.req_sampling_params_manager.req_to_next_token_ids, - b_req_idx=model_input1.b_req_idx, - b_mtp_index=model_input1.b_mtp_index, - b_is_decode_req=model_input1.b_is_decode_req, - b_prefill_start_loc=model_input1.b_prefill_start_loc, - ) + """执行由调用方提前构建好的两个 prefill microbatch。""" + + for model_input in (model_input0, model_input1): + model_input.to_cuda() + if self.args.enable_prefill_decode_mixed and model_input.input_ids.shape[0] > 0: + gather_token_prefill_decode_mixed( + input_ids=model_input.input_ids, + req_to_next_token_ids=(self.req_manager.req_sampling_params_manager.req_to_next_token_ids), + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, + b_is_decode_req=model_input.b_is_decode_req, + b_prefill_start_loc=model_input.b_prefill_start_loc, + ) + return self._microbatch_overlap_prefill_cuda(model_input0, model_input1) + def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input1: ModelInput): assert model_input0.mem_indexes.is_cuda assert model_input1.mem_indexes.is_cuda assert self.args.enable_tpsp_mix_mode origin_handle_token_num0 = model_input0.input_ids.shape[0] origin_handle_token_num1 = model_input1.input_ids.shape[0] - infer_handle_token_num0 = triton.cdiv(origin_handle_token_num0, self.tp_world_size_) * self.tp_world_size_ - infer_handle_token_num1 = triton.cdiv(origin_handle_token_num1, self.tp_world_size_) * self.tp_world_size_ + infer_handle_token_num0 = max( + self.tp_world_size_, + triton.cdiv(origin_handle_token_num0, self.tp_world_size_) * self.tp_world_size_, + ) + infer_handle_token_num1 = max( + self.tp_world_size_, + triton.cdiv(origin_handle_token_num1, self.tp_world_size_) * self.tp_world_size_, + ) origin_batch_size0 = model_input0.batch_size origin_batch_size1 = model_input1.batch_size @@ -867,36 +866,39 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod @torch.no_grad() def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: ModelInput): - model_input0.to_cuda() - model_input1.to_cuda() + """执行由调用方提前构建好的两个 decode microbatch。""" + + for model_input in (model_input0, model_input1): + model_input.to_cuda() + if model_input.input_ids is None: + if model_input.batch_size > 0: + model_input.input_ids = gather_token( + req_to_next_token_ids=(self.req_manager.req_sampling_params_manager.req_to_next_token_ids), + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, + ) + else: + model_input.input_ids = torch.empty( + (0,), + dtype=torch.int64, + device=model_input.b_req_idx.device, + ) + return self._microbatch_overlap_decode_cuda(model_input0, model_input1) + + def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1: ModelInput): assert self.args.enable_tpsp_mix_mode - - if model_input0.input_ids is None: - model_input0.input_ids = gather_token( - self.req_manager.req_sampling_params_manager.req_to_next_token_ids, - model_input0.b_req_idx, - model_input0.b_mtp_index, - ) - if model_input1.input_ids is None: - model_input1.input_ids = gather_token( - self.req_manager.req_sampling_params_manager.req_to_next_token_ids, - model_input1.b_req_idx, - model_input1.b_mtp_index, - ) - # TODO 动态 mtp fix - assert model_input0.batch_size == model_input1.batch_size assert model_input0.mem_indexes.is_cuda assert model_input1.mem_indexes.is_cuda - origin_batch_size = model_input0.batch_size - max_len_in_batch = max(model_input0.max_kv_seq_len, model_input1.max_kv_seq_len) - infer_batch_size = triton.cdiv(origin_batch_size, self.tp_world_size_) * self.tp_world_size_ + origin_batch_size0 = model_input0.batch_size + origin_batch_size1 = model_input1.batch_size + max_len_in_batch = max(2, model_input0.max_kv_seq_len, model_input1.max_kv_seq_len) + infer_batch_size = max(1, origin_batch_size0, origin_batch_size1) + infer_batch_size = triton.cdiv(infer_batch_size, self.tp_world_size_) * self.tp_world_size_ if self.graph is not None and self.graph.can_run(infer_batch_size, max_len_in_batch): infer_batch_size = self.graph.find_closest_graph_batch_size(infer_batch_size) need_capture = self.graph.need_capture(infer_batch_size) - # TODO 如果支持动态步数的 mtp,在不同的mtp步上,model_input0 和 model_input1 的内部batch size可能不 - # 一致,需要按照较高 batch size 进行graph的寻找,同时,进行有效的恢复。 padded_model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) padded_model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) infer_state0 = self._create_inferstate(padded_model_input0, 0) @@ -933,9 +935,8 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state1=infer_state1, ) - # TODO 动态 mtp fix - model_output0 = self._create_unpad_decode_model_output(model_output0, origin_batch_size=origin_batch_size) - model_output1 = self._create_unpad_decode_model_output(model_output1, origin_batch_size=origin_batch_size) + model_output0 = self._create_unpad_decode_model_output(model_output0, origin_batch_size=origin_batch_size0) + model_output1 = self._create_unpad_decode_model_output(model_output1, origin_batch_size=origin_batch_size1) else: model_input0 = self._create_padded_decode_model_input(model_input0, infer_batch_size) model_input1 = self._create_padded_decode_model_input(model_input1, infer_batch_size) @@ -960,8 +961,8 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state1.init_att_state() model_output0, model_output1 = self._overlap_tpsp_token_forward(infer_state0, infer_state1=infer_state1) - model_output0 = self._create_unpad_decode_model_output(model_output0, origin_batch_size=origin_batch_size) - model_output1 = self._create_unpad_decode_model_output(model_output1, origin_batch_size=origin_batch_size) + model_output0 = self._create_unpad_decode_model_output(model_output0, origin_batch_size=origin_batch_size0) + model_output1 = self._create_unpad_decode_model_output(model_output1, origin_batch_size=origin_batch_size1) return model_output0, model_output1 diff --git a/lightllm/common/basemodel/triton_kernel/gather_token_id.py b/lightllm/common/basemodel/triton_kernel/gather_token_id.py index 16c7528b33..a5dfc6d226 100644 --- a/lightllm/common/basemodel/triton_kernel/gather_token_id.py +++ b/lightllm/common/basemodel/triton_kernel/gather_token_id.py @@ -229,7 +229,11 @@ def test_gather_token(): req_ids = torch.arange(20, 20 + batch_size, dtype=torch.int32).cuda() mtp_index = torch.zeros((batch_size,), dtype=torch.int32).cuda() scatter_token(token_info, req_to_token_info, req_ids, mtp_index) - output = gather_token(req_to_token_info, req_ids, mtp_index) + output = gather_token( + req_to_next_token_ids=req_to_token_info, + b_req_idx=req_ids, + b_mtp_index=mtp_index, + ) diff = (token_info - output).abs().max() assert diff < 1e-6 print("test_gather_token passed") diff --git a/lightllm/common/basemodel/triton_kernel/gen_mtp_prefill_params.py b/lightllm/common/basemodel/triton_kernel/gen_mtp_prefill_params.py index 8b1ca912f6..95b92f0de3 100644 --- a/lightllm/common/basemodel/triton_kernel/gen_mtp_prefill_params.py +++ b/lightllm/common/basemodel/triton_kernel/gen_mtp_prefill_params.py @@ -32,6 +32,10 @@ def gen_mtp_new_input_ids( ): assert len(b_seq_len.shape) == 1 batch_size = b_seq_len.shape[0] + if batch_size == 0: + # Overlap prefill 允许一侧保持真实的 0-shape ModelInput。空侧没有 + # token 需要移动,直接返回空输出,避免启动 0-grid Triton kernel。 + return torch.empty_like(input_ids) if b_ready_cache_len is None: b_q_seq_len = b_seq_len else: diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 9b011aad6a..9c99eed04b 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -8,10 +8,8 @@ from lightllm.server.router.model_infer.mode_backend.pre import ( prepare_prefill_inputs, prepare_decode_inputs, - padded_prepare_prefill_inputs, - padded_prepare_decode_inputs, - padded_overlap_prepare_prefill_inputs, - padded_overlap_prepare_decode_inputs, + overlap_prepare_prefill_inputs, + overlap_prepare_decode_inputs, ) from lightllm.server.router.model_infer.mode_backend.overlap_events import OverlapEventPack from lightllm.utils.dist_utils import get_current_device_id @@ -304,11 +302,9 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer ( model_input0, run_reqs0, - _, model_input1, run_reqs1, - _, - ) = padded_overlap_prepare_prefill_inputs(prefill_reqs) + ) = overlap_prepare_prefill_inputs(prefill_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output0, model_output1 = self.model.microbatch_overlap_prefill(model_input0, model_input1) @@ -318,19 +314,34 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer logits1 = model_output1.logits req_num0, req_num1 = len(run_reqs0), len(run_reqs1) - logits = torch.empty((req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) - + logits = torch.empty( + (req_num0 + req_num1, logits0.shape[1]), + dtype=logits0.dtype, + device=logits0.device, + ) logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) - logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1[0:req_num1, :], non_blocking=True) + logits[req_num0 : req_num0 + req_num1, :].copy_(logits1[0:req_num1, :], non_blocking=True) run_reqs = run_reqs0 + run_reqs1 b_has_out_cpu = ( model_input0.b_prefill_has_output_cpu[0:req_num0] + model_input1.b_prefill_has_output_cpu[0:req_num1] ) - b_mtp_index = torch.cat((model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0) - b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) + b_mtp_index = torch.cat( + ( + model_input0.b_mtp_index[0:req_num0], + model_input1.b_mtp_index[0:req_num1], + ), + dim=0, + ) + b_req_idx = torch.cat( + ( + model_input0.b_req_idx[0:req_num0], + model_input1.b_req_idx[0:req_num1], + ), + dim=0, + ) - if (req_num0 + req_num1) > 0: + if req_num0 + req_num1 > 0: ( _, next_token_ids_cpu, @@ -352,7 +363,7 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer sync_event = torch.cuda.Event() sync_event.record() - if (req_num0 + req_num1) > 0: + if req_num0 + req_num1 > 0: # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() update_packs = self._pre_post_handle(run_reqs, is_chuncked_mode=not self.disable_chunked_prefill) @@ -379,32 +390,16 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer return def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): - ( - model_input0, - run_reqs0, - _, - model_input1, - run_reqs1, - _, - ) = padded_overlap_prepare_decode_inputs(req_objs=decode_reqs) - model_input0: ModelInput = model_input0 - model_input1: ModelInput = model_input1 + model_input0, run_reqs0, model_input1, run_reqs1 = overlap_prepare_decode_inputs(req_objs=decode_reqs) + run_reqs = run_reqs0 + run_reqs1 + req_num0, req_num1 = len(run_reqs0), len(run_reqs1) with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output0, model_output1 = self.model.microbatch_overlap_decode(model_input0, model_input1) - logits0 = model_output0.logits - logits1 = model_output1.logits - - req_num0, req_num1 = len(run_reqs0), len(run_reqs1) - logits = torch.empty((req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) - - logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) - logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1[0:req_num1, :], non_blocking=True) - b_mtp_index = torch.cat((model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0) - b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) - - run_reqs = run_reqs0 + run_reqs1 - if (req_num0 + req_num1) > 0: + if req_num0 + req_num1 > 0: + logits = torch.cat((model_output0.logits, model_output1.logits), dim=0) + b_req_idx = torch.cat((model_input0.b_req_idx, model_input1.b_req_idx), dim=0) + b_mtp_index = torch.cat((model_input0.b_mtp_index, model_input1.b_mtp_index), dim=0) ( _, next_token_ids_cpu, @@ -421,7 +416,7 @@ def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferRe sync_event = torch.cuda.Event() sync_event.record() - if (req_num0 + req_num1) > 0: + if req_num0 + req_num1 > 0: # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() update_packs = self._pre_post_handle(run_reqs, is_chuncked_mode=False) @@ -448,7 +443,10 @@ def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferRe def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): # main model prefill - model_input, run_reqs, _ = padded_prepare_prefill_inputs(prefill_reqs) + model_input, run_reqs = prepare_prefill_inputs( + prefill_reqs, + is_chuncked_mode=not self.disable_chunked_prefill, + ) req_num = len(run_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output: ModelOutput = self.model.forward(model_input) @@ -457,6 +455,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] b_req_idx = model_input.b_req_idx[0:req_num] b_mtp_index = model_input.b_mtp_index[0:req_num] + next_token_ids = torch.empty((0,), dtype=torch.int64, device=model_input.b_req_idx.device) if req_num > 0: ( next_token_ids, @@ -520,7 +519,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] return def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): - model_input, run_reqs, _ = padded_prepare_decode_inputs(decode_reqs) + model_input, run_reqs = prepare_decode_inputs(req_objs=decode_reqs) b_mtp_index_cpu = model_input.b_mtp_index req_num = len(run_reqs) @@ -733,11 +732,9 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I ( model_input0, run_reqs0, - _, model_input1, run_reqs1, - _, - ) = padded_overlap_prepare_prefill_inputs(prefill_reqs) + ) = overlap_prepare_prefill_inputs(prefill_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output0, model_output1 = self.model.microbatch_overlap_prefill(model_input0, model_input1) self._capture_prompt_logprobs_if_needed(model_input0, run_reqs0, model_output0.prompt_logics) @@ -830,16 +827,13 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf ( model_input0, run_reqs0, - _, model_input1, run_reqs1, - _, - ) = padded_overlap_prepare_decode_inputs(decode_reqs) + ) = overlap_prepare_decode_inputs(req_objs=decode_reqs) req_num0, req_num1 = len(run_reqs0), len(run_reqs1) b_mtp_index_cpu0 = model_input0.b_mtp_index b_mtp_index_cpu1 = model_input1.b_mtp_index with torch.cuda.stream(g_infer_context.get_overlap_stream()): - model_output0, model_output1 = self.model.microbatch_overlap_decode(model_input0, model_input1) logits0 = model_output0.logits logits1 = model_output1.logits diff --git a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py deleted file mode 100644 index 0e4a562348..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ /dev/null @@ -1,270 +0,0 @@ -import torch -import torch.distributed as dist -import torch.nn.functional as F -import numpy as np -import triton -from typing import List, Optional, Tuple -from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq -from lightllm.utils.infer_utils import calculate_time -from lightllm.utils.envs_utils import get_env_start_args -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from .generic_pre_process import ( - INT64_MAX, - build_b_position_delta, -) - - -def padded_prepare_prefill_inputs( - req_objs: List[InferReq], dest_batch_size: Optional[int] = None -) -> Tuple[ModelInput, List[InferReq], int]: - - if dest_batch_size is None: - req_num = len(req_objs) - if req_num > 0: - dest_batch_size = req_num - else: - dest_batch_size = 1 - else: - assert len(req_objs) <= dest_batch_size - - run_reqs = [] - total_token_num = 0 - padded_req_num = dest_batch_size - len(req_objs) - input_ids = [] - b_req_idx = [] - b_seq_len = [] - b_q_seq_len = [] - batch_multimodal_params = [] - b_ready_cache_len = [] - b_mtp_index = [] - b_prefill_has_output = [] - b_is_decode_req = [] - - for req in req_objs: - - run_reqs.append(req) - batch_multimodal_params.append(req.multimodal_params) - b_req_idx.append(req.req_idx) - - input_token_ids = req.get_chuncked_input_token_ids() - b_prefill_has_output.append(False if len(input_token_ids) < req.get_cur_total_len() else True) - seq_len = len(input_token_ids) - input_token_len = seq_len - req.cur_kv_len - input_id = input_token_ids[req.cur_kv_len :] - - b_seq_len.append(seq_len) - b_q_seq_len.append(input_token_len) - input_ids.append(input_id) - total_token_num += seq_len - b_ready_cache_len.append(req.cur_kv_len) - b_mtp_index.append(0) - - # enable_prefill_decode_mixed 模式下,decode 请求混合在 prefill 请求中。 - # 需要的特殊标记。 - if hasattr(req, "is_decode_req_mixed_in_prefill"): - b_is_decode_req.append(True) - del req.is_decode_req_mixed_in_prefill - else: - b_is_decode_req.append(False) - - # padding fake req for prefill - for _ in range(padded_req_num): - input_ids.append([1]) - b_req_idx.append(g_infer_context.req_manager.HOLD_REQUEST_ID) - b_seq_len.append(1) - b_q_seq_len.append(1) - b_mtp_index.append(0) - b_prefill_has_output.append(False) - b_ready_cache_len.append(0) - total_token_num += 1 - batch_multimodal_params.append({"images": [], "audios": []}) - b_is_decode_req.append(False) - - max_kv_seq_len = max(b_seq_len) - max_cache_len = max(b_ready_cache_len) - max_q_seq_len = max(b_q_seq_len) - - input_ids = np.concatenate(input_ids, dtype=np.int64) - input_ids = torch.tensor(input_ids, dtype=torch.int64, device="cpu") - b_req_idx = torch.tensor(b_req_idx, dtype=torch.int32, device="cpu") - b_seq_len = torch.tensor(b_seq_len, dtype=torch.int32, device="cpu") - b_is_decode_req = torch.tensor(b_is_decode_req, dtype=torch.bool, device="cpu") - b_mtp_index = torch.tensor(b_mtp_index, dtype=torch.int32, device="cpu") - b_ready_cache_len = torch.tensor(b_ready_cache_len, dtype=torch.int32, device="cpu") - b_q_seq_len = torch.tensor(b_q_seq_len, dtype=torch.int32, device="cpu") - b_prefill_start_loc = b_q_seq_len.cumsum(dim=0, dtype=torch.int32) - b_q_seq_len - - # dynamic prompt cache 准备 token - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(input_ids.shape[0] - padded_req_num) - mem_indexes = g_infer_context.req_manager.mem_manager.alloc(input_ids.shape[0] - padded_req_num) - - if padded_req_num > 0: - mem_indexes = F.pad( - input=mem_indexes, - pad=(0, padded_req_num), - mode="constant", - value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, - ) - - model_input = ModelInput( - batch_size=b_seq_len.shape[0], - total_token_num=total_token_num, - max_q_seq_len=max_q_seq_len, - max_kv_seq_len=max_kv_seq_len, - max_cache_len=max_cache_len, - input_ids=input_ids, - mem_indexes_cpu=mem_indexes, - b_req_idx=b_req_idx, - b_mtp_index=b_mtp_index, - b_seq_len=b_seq_len, - b_is_decode_req=b_is_decode_req, - b_ready_cache_len=b_ready_cache_len, - b_prefill_start_loc=b_prefill_start_loc, - is_prefill=True, - b_prefill_has_output_cpu=b_prefill_has_output, - multimodal_params=batch_multimodal_params, - ) - - return model_input, run_reqs, padded_req_num - - -def padded_prepare_decode_inputs( - req_objs: List[InferReq], dest_batch_size: Optional[int] = None -) -> Tuple[ModelInput, List[InferReq], int]: - - if dest_batch_size is None: - if len(req_objs) == 0: - dest_batch_size = 1 - else: - dest_batch_size = len(req_objs) - else: - assert len(req_objs) <= dest_batch_size - - padded_req_num = dest_batch_size - len(req_objs) - - run_reqs = [] - total_token_num = 0 - b_req_idx = [] - b_mtp_index = [] - b_seq_len = [] - b_q_seq_len = [] - args_mtp_step = get_env_start_args().mtp_step - batch_multimodal_params = [] - for req in req_objs: - run_reqs.append(req) - b_req_idx.append(req.req_idx) - seq_len = req.get_cur_total_len() - assert req.cur_kv_len == seq_len - 1 - b_seq_len.append(seq_len) - b_q_seq_len.append(1) - total_token_num += seq_len - b_mtp_index.append(0) - batch_multimodal_params.append(req.multimodal_params) - # process the draft tokens. - for step in range(req.mtp_step): - run_reqs.append(req) - seq_len += 1 - total_token_num += seq_len - b_req_idx.append(req.req_idx) - b_seq_len.append(seq_len) - b_q_seq_len.append(1) - b_mtp_index.append(step + 1) - batch_multimodal_params.append(req.multimodal_params) - - # padding fake req for decode - for _ in range(padded_req_num): - seq_len = 2 - total_token_num += seq_len - b_req_idx.append(g_infer_context.req_manager.HOLD_REQUEST_ID) - b_seq_len.append(seq_len) - b_q_seq_len.append(1) - b_mtp_index.append(0) - batch_multimodal_params.append({"images": [], "audios": []}) - for step in range(args_mtp_step): - seq_len += 1 - total_token_num += seq_len - b_seq_len.append(seq_len) - b_q_seq_len.append(1) - b_req_idx.append(g_infer_context.req_manager.HOLD_REQUEST_ID) - b_mtp_index.append(step + 1) - batch_multimodal_params.append({"images": [], "audios": []}) - - max_kv_seq_len = max(b_seq_len) - max_q_seq_len = max(b_q_seq_len) - - b_req_idx = torch.tensor(b_req_idx, dtype=torch.int32, device="cpu") - b_seq_len = torch.tensor(b_seq_len, dtype=torch.int32, device="cpu") - b_mtp_index = torch.tensor(b_mtp_index, dtype=torch.int32, device="cpu") - b_position_delta = build_b_position_delta(batch_multimodal_params) - - padded_row_count = padded_req_num * (args_mtp_step + 1) - b_shared_seq_len = torch.tensor( - [req.get_radix_cache_shared_len() for req in run_reqs], dtype=torch.int32, device="cpu" - ) - b_shared_radix_node_id = torch.tensor( - [-1 if req.shared_kv_node is None else req.shared_kv_node.time_id % INT64_MAX for req in run_reqs], - dtype=torch.int64, - device="cpu", - ) - if padded_row_count > 0: - b_shared_seq_len = F.pad(b_shared_seq_len, (0, padded_row_count), value=0) - b_shared_radix_node_id = F.pad(b_shared_radix_node_id, (0, padded_row_count), value=-1) - - # dynamic prompt cache 准备 token - if g_infer_context.radix_cache is not None: - g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(b_seq_len.shape[0] - padded_row_count) - mem_indexes = g_infer_context.req_manager.mem_manager.alloc(b_seq_len.shape[0] - padded_row_count) - - if padded_row_count > 0: - mem_indexes = F.pad( - input=mem_indexes, - pad=(0, padded_row_count), - mode="constant", - value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, - ) - - model_input = ModelInput( - batch_size=b_seq_len.shape[0], - total_token_num=total_token_num, - max_q_seq_len=max_q_seq_len, - max_kv_seq_len=max_kv_seq_len, - input_ids=None, - mem_indexes_cpu=mem_indexes, - b_req_idx=b_req_idx, - b_mtp_index=b_mtp_index, - b_seq_len=b_seq_len, - b_position_delta=b_position_delta, - b_shared_seq_len=b_shared_seq_len, - b_shared_radix_node_id=b_shared_radix_node_id, - is_prefill=False, - multimodal_params=batch_multimodal_params, - ) - return model_input, run_reqs, padded_req_num - - -def padded_overlap_prepare_decode_inputs( - req_objs: List[InferReq], -) -> Tuple[ModelInput, List[InferReq], int, ModelInput, List[InferReq], int]: - split_req_bound = triton.cdiv(len(req_objs), 2) - req_objs_0 = req_objs[0:split_req_bound] - req_objs_1 = req_objs[split_req_bound:] - - micro_batch_size = triton.cdiv(len(req_objs), 2) - micro_batch_size = max(1, micro_batch_size) - - micro_input, run_reqs, padded_req_num = padded_prepare_decode_inputs(req_objs_0, dest_batch_size=micro_batch_size) - micro_input1, run_reqs1, padded_req_num1 = padded_prepare_decode_inputs( - req_objs_1, dest_batch_size=micro_batch_size - ) - return micro_input, run_reqs, padded_req_num, micro_input1, run_reqs1, padded_req_num1 - - -def padded_overlap_prepare_prefill_inputs(req_objs: List[InferReq]): - micro_batch1_req_num = triton.cdiv(len(req_objs), 2) - - micro_input, run_reqs, padded_req_num = padded_prepare_prefill_inputs(req_objs[0:micro_batch1_req_num]) - - micro_input1, run_reqs1, padded_req_num1 = padded_prepare_prefill_inputs(req_objs[micro_batch1_req_num:]) - - return micro_input, run_reqs, padded_req_num, micro_input1, run_reqs1, padded_req_num1 diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 313aea64f1..5e73e11f58 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -164,6 +164,54 @@ def prepare_decode_inputs(req_objs: List[InferReq]) -> Tuple[ModelInput, List[In return model_input, run_reqs +def overlap_prepare_decode_inputs(req_objs: List[InferReq]): + """按请求把 decode batch 拆成两个允许为空的 microbatch。""" + + split_req_bound = (len(req_objs) + 1) // 2 + model_input0, run_reqs0 = prepare_decode_inputs( + req_objs=req_objs[:split_req_bound], + ) + model_input1, run_reqs1 = prepare_decode_inputs( + req_objs=req_objs[split_req_bound:], + ) + return model_input0, run_reqs0, model_input1, run_reqs1 + + +def overlap_prepare_prefill_inputs(req_objs: List[InferReq]): + """按当前 prefill token 负载把完整请求分配到两个 microbatch。 + + 请求不能跨 microbatch 拆分,否则同一个 ``InferReq`` 会在后处理阶段被 + 重复更新。所有请求统一按照两侧已分配 token 数进行贪心均衡。这里不 + 创建 HOLD 请求,空侧保留完整的 0-shape ``ModelInput``,执行阶段需要 + 的 padding 由 BaseModel 统一处理。 + """ + + req_input_token_nums = [len(req.get_chuncked_input_token_ids()) - req.cur_kv_len for req in req_objs] + assert all(token_num > 0 for token_num in req_input_token_nums) + + left_token_num = 0 + right_token_num = 0 + left_reqs = [] + right_reqs = [] + for req, token_num in zip(req_objs, req_input_token_nums): + if left_token_num <= right_token_num: + left_reqs.append(req) + left_token_num += token_num + else: + right_reqs.append(req) + right_token_num += token_num + + model_input0, run_reqs0 = prepare_prefill_inputs( + req_objs=left_reqs, + is_chuncked_mode=True, + ) + model_input1, run_reqs1 = prepare_prefill_inputs( + req_objs=right_reqs, + is_chuncked_mode=True, + ) + return model_input0, run_reqs0, model_input1, run_reqs1 + + def build_b_position_delta(multimodal_params: List[dict]) -> torch.Tensor: b_position_delta = [] for params in multimodal_params: diff --git a/lightllm/server/router/model_infer/mode_backend/pre.py b/lightllm/server/router/model_infer/mode_backend/pre.py index e88b1a62b5..7cf9c7c7ae 100644 --- a/lightllm/server/router/model_infer/mode_backend/pre.py +++ b/lightllm/server/router/model_infer/mode_backend/pre.py @@ -1,6 +1,6 @@ -from .generic_pre_process import prepare_prefill_inputs -from .generic_pre_process import prepare_decode_inputs -from .generic_padded_pre_process import padded_prepare_prefill_inputs -from .generic_padded_pre_process import padded_prepare_decode_inputs -from .generic_padded_pre_process import padded_overlap_prepare_prefill_inputs -from .generic_padded_pre_process import padded_overlap_prepare_decode_inputs +from .generic_pre_process import ( + overlap_prepare_decode_inputs, + overlap_prepare_prefill_inputs, + prepare_decode_inputs, + prepare_prefill_inputs, +) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index 3c7ad489fe..0921422fc2 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -48,14 +48,14 @@ def fill_draft_model_kv_state_overlap( def propose_next_overlap( self, - target_model_input0: ModelInput, # batch_size = padded_verify_batch_size0 - target_model_output0: ModelOutput, # logits: [padded_verify_batch_size0, vocab_size] - target_next_token_ids0: torch.Tensor, # [padded_verify_batch_size0] + target_model_input0: ModelInput, # batch_size = verify_batch_size0 + target_model_output0: ModelOutput, # logits: [verify_batch_size0, vocab_size] + target_next_token_ids0: torch.Tensor, # [verify_batch_size0] real_verify_rows0: int, accept_len0: Optional[torch.Tensor], # [req_num0] - target_model_input1: ModelInput, # batch_size = padded_verify_batch_size1 - target_model_output1: ModelOutput, # logits: [padded_verify_batch_size1, vocab_size] - target_next_token_ids1: torch.Tensor, # [padded_verify_batch_size1] + target_model_input1: ModelInput, # batch_size = verify_batch_size1 + target_model_output1: ModelOutput, # logits: [verify_batch_size1, vocab_size] + target_next_token_ids1: torch.Tensor, # [verify_batch_size1] real_verify_rows1: int, accept_len1: Optional[torch.Tensor], # [req_num1] ) -> SpecProposal: diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index e276540f67..27af2c5dfb 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -41,7 +41,7 @@ def fill_dp_eagle_draft_model_kv_state_overlap( b_next_token_ids=target_next_token_ids1, mtp_draft_input_hiddens=target_model_output1.mtp_collector.spec_hidden, ) - proposer.backend.draft_models[0].microbatch_overlap_prefill(target_model_input0, target_model_input1) + proposer.backend.draft_models[0]._microbatch_overlap_prefill_cuda(target_model_input0, target_model_input1) def pad_dp_step_mem_indexes( @@ -118,7 +118,7 @@ def propose_next_dp_eagle_autoregressive_overlap( proposal_token_ids = target_next_token_ids0.new_empty((total_real_request_count, draft_step)) draft_model = proposer.backend.draft_models[0] - extend_outputs = draft_model.microbatch_overlap_decode(*model_inputs) + extend_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) draft_token_ids_by_batch = [] draft_hiddens_by_batch = [] @@ -190,7 +190,7 @@ def propose_next_dp_eagle_autoregressive_overlap( model_input.max_kv_seq_len = max_kv_seq_lens_by_batch[batch_index] + step model_input.total_token_num = model_input.batch_size * model_input.max_kv_seq_len - draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) + draft_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) for batch_index, draft_output in enumerate(draft_outputs): draft_token_ids = generate_eagle_token_ids(proposer, draft_output, map_draft_token_ids) draft_token_ids_by_batch[batch_index] = draft_token_ids @@ -265,7 +265,7 @@ def propose_next_dp_eagle_fixed_layout_overlap( model_input.input_ids = draft_token_ids_by_batch[batch_index] model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - draft_outputs = draft_model.microbatch_overlap_decode(*model_inputs) + draft_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) for batch_index, (model_input, draft_output) in enumerate(zip(model_inputs, draft_outputs)): model_input.b_seq_len += 1 diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 5a9961957e..2489dd9d26 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -124,7 +124,7 @@ def propose_next_overlap( model_input.input_ids = draft_token_ids[batch_index] model_input.mtp_draft_input_hiddens = draft_hiddens[batch_index] - draft_outputs = self.backend.draft_models[step].microbatch_overlap_decode(*model_inputs) + draft_outputs = self.backend.draft_models[step]._microbatch_overlap_decode_cuda(*model_inputs) for batch_index, draft_output in enumerate(draft_outputs): draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index c95a335bcc..8bbfcd7214 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -66,7 +66,7 @@ def fill_draft_model_kv_state_overlap( b_next_token_ids=draft_token_ids[batch_index], mtp_draft_input_hiddens=draft_hiddens[batch_index], ) - draft_outputs = draft_model.microbatch_overlap_prefill(*model_inputs) + draft_outputs = draft_model._microbatch_overlap_prefill_cuda(*model_inputs) for batch_index, draft_output in enumerate(draft_outputs): draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) @@ -171,7 +171,7 @@ def propose_next_overlap( model_input.input_ids = draft_token_ids[batch_index] model_input.mtp_draft_input_hiddens = draft_hiddens[batch_index] - draft_outputs = self.backend.draft_models[step].microbatch_overlap_decode(*model_inputs) + draft_outputs = self.backend.draft_models[step]._microbatch_overlap_decode_cuda(*model_inputs) for batch_index, draft_output in enumerate(draft_outputs): draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) diff --git a/test/benchmark/static_inference/static_benchmark.py b/test/benchmark/static_inference/static_benchmark.py index b3f19f360a..8c0a280d78 100644 --- a/test/benchmark/static_inference/static_benchmark.py +++ b/test/benchmark/static_inference/static_benchmark.py @@ -26,7 +26,7 @@ if str(REPO_ROOT) not in sys.path: sys.path.append(str(REPO_ROOT)) -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.batch_objs import ModelInput, ModelMtpOutputCollector, ModelOutput from lightllm.models import get_model from lightllm.models.deepseek_mtp.model import Deepseek3MTPModel from lightllm.models.glm4_moe_lite_mtp.model import Glm4MoeLiteMTPModel @@ -572,6 +572,8 @@ def _make_decode_input( b_mtp_index=mtp_index, b_seq_len=seq_len, b_position_delta=cpu_i32_zeros(batch_size), + b_shared_seq_len=cpu_i32_zeros(batch_size), + b_shared_radix_node_id=torch.full((batch_size,), -1, dtype=torch.int64, device="cpu"), mem_indexes_cpu=mem_indexes, is_prefill=False, multimodal_params=empty_multimodal_params(batch_size), @@ -579,19 +581,21 @@ def _make_decode_input( def _forward_prefill_input(self, model_input: ModelInput, allow_overlap: bool) -> ModelOutput: if allow_overlap and self.args.enable_prefill_microbatch_overlap and model_input.batch_size > 1: - micro_input0, micro_input1 = self._split_prefill_input(model_input) - output0, output1 = self.model.microbatch_overlap_prefill(micro_input0, micro_input1) - return self._merge_model_outputs(output0, output1) + model_input0, model_input1 = self._split_prefill_input(model_input) + model_output0, model_output1 = self.model.microbatch_overlap_prefill(model_input0, model_input1) + return self._merge_overlap_model_outputs(model_output0, model_output1) return self.model.forward(model_input) def _forward_decode_input(self, model_input: ModelInput, allow_overlap: bool) -> ModelOutput: - if allow_overlap and self.args.enable_decode_microbatch_overlap and model_input.batch_size > 1: - micro_input0, micro_input1 = self._split_decode_input(model_input) - output0, output1 = self.model.microbatch_overlap_decode(micro_input0, micro_input1) - return self._merge_model_outputs(output0, output1) + if allow_overlap and self.args.enable_decode_microbatch_overlap: + model_input0, model_input1 = self._split_decode_input(model_input) + model_output0, model_output1 = self.model.microbatch_overlap_decode(model_input0, model_input1) + return self._merge_overlap_model_outputs(model_output0, model_output1) return self.model.forward(model_input) def _split_prefill_input(self, model_input: ModelInput): + """在 benchmark 调用侧按请求构造两个 prefill microbatch。""" + split_batch = model_input.batch_size // 2 q_lens = model_input.b_seq_len - model_input.b_ready_cache_len split_tokens = int(q_lens[:split_batch].sum().item()) @@ -616,65 +620,101 @@ def _slice_prefill_input( ) -> ModelInput: b_seq_len = model_input.b_seq_len[batch_start:batch_end].clone() b_ready_cache_len = model_input.b_ready_cache_len[batch_start:batch_end].clone() - b_q_seq_len = b_seq_len - b_ready_cache_len - b_prefill_start_loc = b_q_seq_len.cumsum(dim=0, dtype=torch.int32) - b_q_seq_len + q_lens = b_seq_len - b_ready_cache_len + b_prefill_start_loc = q_lens.cumsum(dim=0, dtype=torch.int32) - q_lens has_output = model_input.b_prefill_has_output_cpu - return ModelInput( + kwargs = dict( batch_size=batch_end - batch_start, total_token_num=int(b_seq_len.sum().item()), - max_q_seq_len=int(b_q_seq_len.max().item()), + max_q_seq_len=int(q_lens.max().item()), max_kv_seq_len=int(b_seq_len.max().item()), max_cache_len=int(b_ready_cache_len.max().item()), input_ids=model_input.input_ids[token_start:token_end].contiguous(), b_req_idx=model_input.b_req_idx[batch_start:batch_end].clone(), b_mtp_index=model_input.b_mtp_index[batch_start:batch_end].clone(), b_seq_len=b_seq_len, - mem_indexes_cpu=model_input.mem_indexes_cpu[token_start:token_end].contiguous(), + b_is_decode_req=model_input.b_is_decode_req[batch_start:batch_end].clone(), is_prefill=True, b_ready_cache_len=b_ready_cache_len, b_prefill_start_loc=b_prefill_start_loc, - b_prefill_has_output_cpu=(has_output[batch_start:batch_end] if has_output is not None else None), + b_prefill_has_output_cpu=has_output[batch_start:batch_end], multimodal_params=model_input.multimodal_params[batch_start:batch_end], ) + if model_input.mem_indexes_cpu is not None: + kwargs["mem_indexes_cpu"] = model_input.mem_indexes_cpu[token_start:token_end].contiguous() + else: + kwargs["mem_indexes"] = model_input.mem_indexes[token_start:token_end].contiguous() + if model_input.mtp_draft_input_hiddens is not None: + kwargs["mtp_draft_input_hiddens"] = model_input.mtp_draft_input_hiddens[token_start:token_end].contiguous() + return ModelInput(**kwargs) def _split_decode_input(self, model_input: ModelInput): - split_batch = model_input.batch_size // 2 + split_batch = (model_input.batch_size + 1) // 2 return ( self._slice_decode_input(model_input, 0, split_batch), self._slice_decode_input(model_input, split_batch, model_input.batch_size), ) - def _slice_decode_input(self, model_input: ModelInput, batch_start: int, batch_end: int) -> ModelInput: + @staticmethod + def _slice_decode_input(model_input: ModelInput, batch_start: int, batch_end: int) -> ModelInput: + batch_size = batch_end - batch_start b_seq_len = model_input.b_seq_len[batch_start:batch_end].clone() input_ids = model_input.input_ids if input_ids is not None: input_ids = input_ids[batch_start:batch_end].contiguous() - return ModelInput( - batch_size=batch_end - batch_start, + kwargs = dict( + batch_size=batch_size, total_token_num=int(b_seq_len.sum().item()), - max_q_seq_len=model_input.max_q_seq_len, - max_kv_seq_len=int(b_seq_len.max().item()), + max_q_seq_len=1, + max_kv_seq_len=int(b_seq_len.max().item()) if batch_size > 0 else 0, input_ids=input_ids, b_req_idx=model_input.b_req_idx[batch_start:batch_end].clone(), b_mtp_index=model_input.b_mtp_index[batch_start:batch_end].clone(), b_seq_len=b_seq_len, b_position_delta=model_input.b_position_delta[batch_start:batch_end].clone(), - mem_indexes_cpu=model_input.mem_indexes_cpu[batch_start:batch_end].contiguous(), + b_shared_seq_len=model_input.b_shared_seq_len[batch_start:batch_end].clone(), + b_shared_radix_node_id=model_input.b_shared_radix_node_id[batch_start:batch_end].clone(), is_prefill=False, multimodal_params=model_input.multimodal_params[batch_start:batch_end], ) + if model_input.mem_indexes_cpu is not None: + kwargs["mem_indexes_cpu"] = model_input.mem_indexes_cpu[batch_start:batch_end].contiguous() + else: + kwargs["mem_indexes"] = model_input.mem_indexes[batch_start:batch_end].contiguous() + if model_input.mtp_draft_input_hiddens is not None: + kwargs["mtp_draft_input_hiddens"] = model_input.mtp_draft_input_hiddens[batch_start:batch_end].contiguous() + return ModelInput(**kwargs) + + @staticmethod + def _merge_overlap_model_outputs(output0: ModelOutput, output1: ModelOutput) -> ModelOutput: + def concat_optional(value0: torch.Tensor | None, value1: torch.Tensor | None): + if value0 is None and value1 is None: + return None + assert value0 is not None and value1 is not None + return torch.cat((value0, value1), dim=0) - def _merge_model_outputs(self, output0: ModelOutput, output1: ModelOutput) -> ModelOutput: - mtp_hiddens = None - if output0.mtp_main_output_hiddens is not None and output1.mtp_main_output_hiddens is not None: - mtp_hiddens = torch.cat( - (output0.mtp_main_output_hiddens, output1.mtp_main_output_hiddens), - dim=0, - ) return ModelOutput( logits=torch.cat((output0.logits, output1.logits), dim=0), - prefill_mem_indexes_ready_event=output0.prefill_mem_indexes_ready_event, - mtp_main_output_hiddens=mtp_hiddens, + mtp_collector=ModelMtpOutputCollector( + spec_hidden=concat_optional( + output0.mtp_collector.spec_hidden, + output1.mtp_collector.spec_hidden, + ), + draft_token_ids=concat_optional( + output0.mtp_collector.draft_token_ids, + output1.mtp_collector.draft_token_ids, + ), + confidence_logits=concat_optional( + output0.mtp_collector.confidence_logits, + output1.mtp_collector.confidence_logits, + ), + ), + prefill_mem_indexes_ready_event=( + output0.prefill_mem_indexes_ready_event + if output0.prefill_mem_indexes_ready_event is not None + else output1.prefill_mem_indexes_ready_event + ), + prompt_logics=concat_optional(output0.prompt_logics, output1.prompt_logics), ) def _build_mtp_decode_index_tensors(self, req_idx: torch.Tensor, step_width: int): diff --git a/unit_tests/common/basemodel/test_overlap_utils.py b/unit_tests/common/basemodel/test_overlap_utils.py new file mode 100644 index 0000000000..b286998e04 --- /dev/null +++ b/unit_tests/common/basemodel/test_overlap_utils.py @@ -0,0 +1,213 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.basemodel import basemodel +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput + + +def _empty_multimodal_params(batch_size: int): + return [{"images": [], "audios": []} for _ in range(batch_size)] + + +def _make_prefill_input(input_ids: list[int], req_idx: int, is_decode_req: bool) -> ModelInput: + token_num = len(input_ids) + return ModelInput( + batch_size=1, + total_token_num=token_num, + max_q_seq_len=token_num, + max_kv_seq_len=token_num, + max_cache_len=0, + input_ids=torch.tensor(input_ids, dtype=torch.int64), + b_req_idx=torch.tensor([req_idx], dtype=torch.int32), + b_mtp_index=torch.zeros(1, dtype=torch.int32), + b_seq_len=torch.tensor([token_num], dtype=torch.int32), + b_is_decode_req=torch.tensor([is_decode_req]), + b_ready_cache_len=torch.zeros(1, dtype=torch.int32), + b_prefill_start_loc=torch.zeros(1, dtype=torch.int32), + mem_indexes_cpu=torch.arange(token_num, dtype=torch.int32), + is_prefill=True, + b_prefill_has_output_cpu=[True], + multimodal_params=_empty_multimodal_params(1), + ) + + +def _make_decode_input() -> ModelInput: + batch_size = 6 + return ModelInput( + batch_size=batch_size, + total_token_num=39, + max_q_seq_len=1, + max_kv_seq_len=9, + input_ids=torch.arange(batch_size, dtype=torch.int64), + b_req_idx=torch.tensor([10, 11, 12, 12, 12, 12], dtype=torch.int32), + b_mtp_index=torch.tensor([0, 0, 0, 1, 2, 3], dtype=torch.int32), + b_seq_len=torch.tensor([4, 5, 6, 7, 8, 9], dtype=torch.int32), + b_position_delta=torch.arange(batch_size, dtype=torch.int32), + b_shared_seq_len=torch.arange(10, 16, dtype=torch.int32), + b_shared_radix_node_id=torch.arange(20, 26, dtype=torch.int64), + mem_indexes_cpu=torch.arange(100, 106, dtype=torch.int32), + is_prefill=False, + multimodal_params=_empty_multimodal_params(batch_size), + ) + + +def _make_empty_decode_input() -> ModelInput: + return ModelInput( + batch_size=0, + total_token_num=0, + max_q_seq_len=1, + max_kv_seq_len=0, + input_ids=torch.empty((0,), dtype=torch.int64), + b_req_idx=torch.empty((0,), dtype=torch.int32), + b_mtp_index=torch.empty((0,), dtype=torch.int32), + b_seq_len=torch.empty((0,), dtype=torch.int32), + b_position_delta=torch.empty((0,), dtype=torch.int32), + b_shared_seq_len=torch.empty((0,), dtype=torch.int32), + b_shared_radix_node_id=torch.empty((0,), dtype=torch.int64), + mem_indexes_cpu=torch.empty((0,), dtype=torch.int32), + is_prefill=False, + multimodal_params=[], + ) + + +def test_base_model_prefill_accepts_two_prebuilt_inputs(monkeypatch): + model_input0 = _make_prefill_input([10, 11], req_idx=20, is_decode_req=False) + model_input1 = _make_prefill_input([12], req_idx=21, is_decode_req=True) + events = [] + original_to_cuda = ModelInput.to_cuda + + def record_to_cuda(self): + events.append("to_cuda") + original_to_cuda(self) + + monkeypatch.setattr(ModelInput, "to_cuda", record_to_cuda) + + def fake_gather(**kwargs): + events.append("gather") + kwargs["input_ids"][-1] = 90 + events.count("gather") + + monkeypatch.setattr(basemodel, "gather_token_prefill_decode_mixed", fake_gather) + + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.args = SimpleNamespace(enable_prefill_decode_mixed=True) + model.req_manager = SimpleNamespace( + req_sampling_params_manager=SimpleNamespace(req_to_next_token_ids=object()), + ) + captured_inputs = [] + + def fake_overlap_forward(input0, input1): + events.append("forward") + captured_inputs.extend((input0, input1)) + return ( + ModelOutput(logits=torch.zeros((input0.batch_size, 1))), + ModelOutput(logits=torch.zeros((input1.batch_size, 1))), + ) + + model._microbatch_overlap_prefill_cuda = fake_overlap_forward + + outputs = model.microbatch_overlap_prefill(model_input0, model_input1) + + assert events == ["to_cuda", "gather", "to_cuda", "gather", "forward"] + assert captured_inputs == [model_input0, model_input1] + assert captured_inputs[0].input_ids.tolist() == [10, 91] + assert captured_inputs[1].input_ids.tolist() == [92] + assert len(outputs) == 2 + + +def test_base_model_decode_accepts_two_inputs_and_skips_empty_gather(monkeypatch): + model_input0 = _make_decode_input() + model_input0.input_ids = None + model_input1 = _make_empty_decode_input() + model_input1.input_ids = None + events = [] + original_to_cuda = ModelInput.to_cuda + + def record_to_cuda(self): + events.append("to_cuda") + original_to_cuda(self) + + monkeypatch.setattr(ModelInput, "to_cuda", record_to_cuda) + + def fake_gather(**kwargs): + events.append("gather") + return torch.arange(40, 46, dtype=torch.int64, device="cuda") + + monkeypatch.setattr(basemodel, "gather_token", fake_gather) + + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.req_manager = SimpleNamespace( + req_sampling_params_manager=SimpleNamespace(req_to_next_token_ids=object()), + ) + captured_inputs = [] + + def fake_overlap_forward(input0, input1): + events.append("forward") + captured_inputs.extend((input0, input1)) + return ( + ModelOutput(logits=torch.zeros((input0.batch_size, 1))), + ModelOutput(logits=torch.zeros((input1.batch_size, 1))), + ) + + model._microbatch_overlap_decode_cuda = fake_overlap_forward + + model.microbatch_overlap_decode(model_input0, model_input1) + + assert events == ["to_cuda", "gather", "to_cuda", "forward"] + assert captured_inputs == [model_input0, model_input1] + assert captured_inputs[0].input_ids.tolist() == [40, 41, 42, 43, 44, 45] + assert captured_inputs[1].input_ids.numel() == 0 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_overlap_decode_cuda_pads_empty_side_and_unpads_outputs(monkeypatch): + model_input0 = _make_decode_input() + model_input0.batch_size = 1 + model_input0.total_token_num = 4 + model_input0.max_kv_seq_len = 4 + model_input0.input_ids = model_input0.input_ids[:1] + model_input0.b_req_idx = model_input0.b_req_idx[:1] + model_input0.b_mtp_index = model_input0.b_mtp_index[:1] + model_input0.b_seq_len = model_input0.b_seq_len[:1] + model_input0.b_position_delta = model_input0.b_position_delta[:1] + model_input0.b_shared_seq_len = model_input0.b_shared_seq_len[:1] + model_input0.b_shared_radix_node_id = model_input0.b_shared_radix_node_id[:1] + model_input0.mem_indexes_cpu = model_input0.mem_indexes_cpu[:1] + model_input0.multimodal_params = model_input0.multimodal_params[:1] + model_input0.check_input() + model_input1 = _make_empty_decode_input() + model_input0.to_cuda() + model_input1.to_cuda() + + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.args = SimpleNamespace(enable_tpsp_mix_mode=True) + model.tp_world_size_ = 2 + model.graph = None + model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) + model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=77) + infer_batch_sizes = [] + + def fake_create_inferstate(model_input, microbatch_index): + infer_batch_sizes.append(model_input.batch_size) + return SimpleNamespace( + b_req_idx=model_input.b_req_idx, + b_seq_len=model_input.b_seq_len, + mem_index=model_input.mem_indexes, + init_some_extra_state=lambda _: None, + init_att_state=lambda: None, + ) + + model._create_inferstate = fake_create_inferstate + model._overlap_tpsp_token_forward = lambda infer_state0, infer_state1: ( + ModelOutput(logits=torch.zeros((infer_state0.b_req_idx.shape[0], 1), device="cuda")), + ModelOutput(logits=torch.zeros((infer_state1.b_req_idx.shape[0], 1), device="cuda")), + ) + monkeypatch.setattr(basemodel, "copy_kv_index_to_req", lambda *args, **kwargs: None) + + output0, output1 = model._microbatch_overlap_decode_cuda(model_input0, model_input1) + + assert infer_batch_sizes == [2, 2] + assert output0.logits.shape == (1, 1) + assert output1.logits.shape == (0, 1) diff --git a/unit_tests/common/basemodel/triton_kernel/test_gen_mtp_prefill_params.py b/unit_tests/common/basemodel/triton_kernel/test_gen_mtp_prefill_params.py index a59af6a31c..2702ca9453 100644 --- a/unit_tests/common/basemodel/triton_kernel/test_gen_mtp_prefill_params.py +++ b/unit_tests/common/basemodel/triton_kernel/test_gen_mtp_prefill_params.py @@ -5,6 +5,22 @@ from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids +def test_gen_mtp_new_input_ids_empty_batch(): + input_ids = torch.empty((0,), dtype=torch.int64, device="cuda") + b_next_token_ids = torch.empty((0,), dtype=torch.int64, device="cuda") + b_seq_len = torch.empty((0,), dtype=torch.int32, device="cuda") + + new_input_ids = gen_mtp_new_input_ids( + input_ids=input_ids, + b_next_token_ids=b_next_token_ids, + b_seq_len=b_seq_len, + ) + + assert new_input_ids.is_cuda + assert new_input_ids.dtype == torch.int64 + assert new_input_ids.shape == (0,) + + def test_gen_mtp_new_input_ids_0(): input_ids = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9]).long().cuda() b_next_token_ids = torch.tensor([10, 11, 12]).long().cuda() diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index ad9d55d2ec..d22c380e19 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -1,7 +1,7 @@ from types import SimpleNamespace import torch -from lightllm.server.router.model_infer.mode_backend import generic_padded_pre_process, generic_pre_process +from lightllm.server.router.model_infer.mode_backend import generic_pre_process def _patch_empty_input_context(monkeypatch): @@ -17,6 +17,34 @@ def _patch_empty_input_context(monkeypatch): return infer_context +def _patch_overlap_input_context(monkeypatch): + return _patch_empty_input_context(monkeypatch) + + +def _make_prefill_req(req_idx: int, token_num: int): + input_token_ids = [req_idx] * token_num + return SimpleNamespace( + req_idx=req_idx, + cur_kv_len=0, + multimodal_params={"images": [], "audios": []}, + get_chuncked_input_token_ids=lambda: input_token_ids, + get_input_token_ids=lambda: input_token_ids, + get_cur_total_len=lambda: token_num, + ) + + +def _make_decode_req(req_idx: int): + return SimpleNamespace( + req_idx=req_idx, + cur_kv_len=3, + mtp_step=0, + multimodal_params={"images": [], "audios": []}, + shared_kv_node=None, + get_cur_total_len=lambda: 4, + get_radix_cache_shared_len=lambda: 0, + ) + + def test_prepare_prefill_inputs_allows_empty_batch(monkeypatch): _patch_empty_input_context(monkeypatch) @@ -50,39 +78,77 @@ def test_prepare_decode_inputs_allows_empty_batch(monkeypatch): assert model_input.max_kv_seq_len == 0 -def test_padded_decode_builds_raw_shared_radix_metadata(monkeypatch): - max_draft_step = 2 - mem_manager = SimpleNamespace( - HOLD_TOKEN_MEMINDEX=-1, - alloc=lambda size: torch.arange(size, dtype=torch.int32), - ) - req_manager = SimpleNamespace(HOLD_REQUEST_ID=-1, mem_manager=mem_manager) - monkeypatch.setattr( - generic_padded_pre_process, - "g_infer_context", - SimpleNamespace(req_manager=req_manager, radix_cache=None), - ) - monkeypatch.setattr( - generic_padded_pre_process, - "get_env_start_args", - lambda: SimpleNamespace(mtp_step=max_draft_step, mtp_dynamic_verify=True), - ) +def test_overlap_prefill_balances_request_token_load_without_padding(monkeypatch): + _patch_overlap_input_context(monkeypatch) + reqs = [ + _make_prefill_req(req_idx=0, token_num=8), + _make_prefill_req(req_idx=1, token_num=7), + _make_prefill_req(req_idx=2, token_num=6), + _make_prefill_req(req_idx=3, token_num=5), + _make_prefill_req(req_idx=4, token_num=1), + ] - shared_kv_node = SimpleNamespace(time_id=torch.iinfo(torch.int64).max + 42) - req = SimpleNamespace( - req_idx=7, - cur_kv_len=4, - mtp_step=max_draft_step, - multimodal_params={"images": [], "audios": []}, - shared_kv_node=shared_kv_node, - get_cur_total_len=lambda: 5, - get_radix_cache_shared_len=lambda: 4, - ) - model_input, _, padded_req_num = generic_padded_pre_process.padded_prepare_decode_inputs( - req_objs=[req], dest_batch_size=3 - ) + ( + model_input0, + run_reqs0, + model_input1, + run_reqs1, + ) = generic_pre_process.overlap_prepare_prefill_inputs(reqs) + + assert [req.req_idx for req in run_reqs0] == [0, 3, 4] + assert [req.req_idx for req in run_reqs1] == [1, 2] + assert model_input0.b_req_idx.tolist() == [0, 3, 4] + assert model_input1.b_req_idx.tolist() == [1, 2] + assert model_input0.input_ids.shape == (14,) + assert model_input1.input_ids.shape == (13,) + assert model_input0.batch_size == 3 + assert model_input1.batch_size == 2 + + +def test_overlap_prefill_balances_single_token_request_normally(monkeypatch): + _patch_overlap_input_context(monkeypatch) + req = _make_prefill_req(req_idx=7, token_num=1) + + ( + model_input0, + run_reqs0, + model_input1, + run_reqs1, + ) = generic_pre_process.overlap_prepare_prefill_inputs([req]) + + assert run_reqs0 == [req] + assert model_input0.batch_size == 1 + assert model_input0.input_ids.tolist() == [7] + assert model_input0.b_req_idx.tolist() == [7] + assert run_reqs1 == [] + assert model_input1.batch_size == 0 + assert model_input1.input_ids.shape == (0,) + assert model_input1.b_req_idx.shape == (0,) + + +def test_overlap_decode_builds_two_unpadded_inputs(monkeypatch): + _patch_overlap_input_context(monkeypatch) + reqs = [_make_decode_req(req_idx=index) for index in range(3)] + + model_input0, run_reqs0, model_input1, run_reqs1 = generic_pre_process.overlap_prepare_decode_inputs(reqs) + + assert run_reqs0 == reqs[:2] + assert run_reqs1 == reqs[2:] + assert model_input0.batch_size == 2 + assert model_input1.batch_size == 1 + assert model_input0.b_req_idx.tolist() == [0, 1] + assert model_input1.b_req_idx.tolist() == [2] + + +def test_overlap_decode_preserves_empty_microbatch(monkeypatch): + _patch_overlap_input_context(monkeypatch) + req = _make_decode_req(req_idx=7) + + model_input0, run_reqs0, model_input1, run_reqs1 = generic_pre_process.overlap_prepare_decode_inputs([req]) - assert padded_req_num == 2 - assert model_input.b_mtp_index.tolist() == [0, 1, 2, 0, 1, 2, 0, 1, 2] - assert model_input.b_shared_seq_len.tolist() == [4, 4, 4, 0, 0, 0, 0, 0, 0] - assert model_input.b_shared_radix_node_id.tolist() == [42, 42, 42, -1, -1, -1, -1, -1, -1] + assert run_reqs0 == [req] + assert model_input0.batch_size == 1 + assert run_reqs1 == [] + assert model_input1.batch_size == 0 + assert model_input1.b_req_idx.shape == (0,) + assert model_input1.mem_indexes_cpu.shape == (0,) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 99c563a3f9..44208139dc 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -18,7 +18,7 @@ def __init__(self): self.decode_batch_sizes = [] self.decode_inputs = [] - def microbatch_overlap_prefill(self, input0, input1): + def _microbatch_overlap_prefill_cuda(self, input0, input1): self.extend_batch_sizes = (input0.batch_size, input1.batch_size) self.extend_inputs = (input0, input1) return tuple( @@ -29,7 +29,7 @@ def microbatch_overlap_prefill(self, input0, input1): for model_input in (input0, input1) ) - def microbatch_overlap_decode(self, input0, input1): + def _microbatch_overlap_decode_cuda(self, input0, input1): self.decode_batch_sizes.append((input0.batch_size, input1.batch_size)) self.decode_inputs.append((input0, input1)) return tuple( diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py index 307ef1f0e9..b6ca60da83 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -13,7 +13,7 @@ class _DraftModel: def __init__(self): self.decode_batch_sizes = [] - def microbatch_overlap_decode(self, input0, input1): + def _microbatch_overlap_decode_cuda(self, input0, input1): self.decode_batch_sizes.append((input0.batch_size, input1.batch_size)) return tuple( ModelOutput( diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py index 040edacc03..c912d81212 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py @@ -100,7 +100,7 @@ class DraftModel: def __init__(self, token_offset): self.token_offset = token_offset - def microbatch_overlap_prefill(self, input0, input1): + def _microbatch_overlap_prefill_cuda(self, input0, input1): forwarded.append((input0, input1, input0.input_ids.clone(), input1.input_ids.clone())) return tuple( SimpleNamespace( From a8659f1340e93bb09be92c191b872e5c674c3d22 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Fri, 21 Aug 2026 16:51:54 +0800 Subject: [PATCH 080/103] fix: account for MTP weights in KV cache profiling (#1478) --- lightllm/common/basemodel/basemodel.py | 11 ++++++-- lightllm/utils/envs_utils.py | 14 ++++++++++ lightllm/utils/profile_max_tokens.py | 36 ++++++++++++++++++++++++++ 3 files changed, 59 insertions(+), 2 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index dbeb15ffde..6ffc3f3535 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -25,7 +25,12 @@ from lightllm.common.basemodel.triton_kernel.gather_token_id import gather_token, gather_token_prefill_decode_mixed from lightllm.utils.log_utils import init_logger from lightllm.utils.dist_utils import get_dp_world_size -from lightllm.utils.envs_utils import get_env_start_args, get_llm_data_type, get_added_mtp_kv_layer_num +from lightllm.utils.profile_max_tokens import profile_mtp_weight_memory +from lightllm.utils.envs_utils import ( + get_env_start_args, + get_llm_data_type, + get_added_mtp_kv_layer_num, +) from lightllm.distributed.communication_op import dist_group_manager from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.hidden_collector import ( @@ -112,7 +117,9 @@ def __init__(self, kvargs): self._init_quant() enable_weight_cpu_backup = self.args.enable_weight_cpu_backup - with self.torch_memory_saver.region(tag=MemoryTag.WEIGHT, enable_cpu_backup=enable_weight_cpu_backup): + with profile_mtp_weight_memory(self), self.torch_memory_saver.region( + tag=MemoryTag.WEIGHT, enable_cpu_backup=enable_weight_cpu_backup + ): self._init_weights() with self.torch_memory_saver.region(tag=MemoryTag.KV_CACHE): self._init_req_manager() diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index b3865b7af2..4fed9509a9 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -274,6 +274,20 @@ def get_added_mtp_kv_layer_num() -> int: raise ValueError(f"unsupported mtp_mode: {mtp_mode}") +@lru_cache(maxsize=None) +def get_mtp_weight_layer_num() -> int: + args = get_env_start_args() + mtp_mode = args.mtp_mode + + if mtp_mode is None: + return 0 + if mtp_mode == "vanilla_no_att": + return args.mtp_step + if mtp_mode == "eagle_no_att": + return 1 + return get_added_mtp_kv_layer_num() + + def _get_mtp_draft_backbone_layer_num(draft_model_dir: str) -> int: with open(os.path.join(draft_model_dir, "config.json"), "r") as json_file: draft_config = json.load(json_file) diff --git a/lightllm/utils/profile_max_tokens.py b/lightllm/utils/profile_max_tokens.py index e3a62b62ea..eb576a966f 100644 --- a/lightllm/utils/profile_max_tokens.py +++ b/lightllm/utils/profile_max_tokens.py @@ -2,10 +2,12 @@ import gc import json import torch +from contextlib import contextmanager from transformers import AutoModelForCausalLM import argparse from lightllm.common.build_utils import repair_config from lightllm.utils.dist_utils import get_current_device_id +from lightllm.utils.envs_utils import get_mtp_weight_layer_num data_type_dict = {"float32": 4, "float16": 2, "bfloat16": 2, "fp32": 4, "fp16": 2, "bf16": 2, "int8": 1, "int4": 0.5} @@ -31,6 +33,40 @@ def get_total_gpu_memory(): return total_memory / (1024 ** 3) # Convert to GB +def get_mtp_adjusted_mem_fraction( + mem_fraction: float, + target_weight_bytes: int, + target_layer_num: int, + mtp_layer_num: int, +) -> float: + mtp_weight_bytes = target_weight_bytes * mtp_layer_num / target_layer_num + total_gpu_bytes = torch.cuda.get_device_properties(get_current_device_id()).total_memory + adjusted_mem_fraction = mem_fraction - mtp_weight_bytes / total_gpu_bytes + assert adjusted_mem_fraction > 0, "MTP weight reservation leaves no memory for the target model KV cache" + return adjusted_mem_fraction + + +@contextmanager +def profile_mtp_weight_memory(model): + """Reserve auto-profiled KV budget for MTP weights allocated after the target model.""" + should_adjust = ( + model.max_total_token_num is None and not model.is_mtp_draft_model and model.args.mtp_mode is not None + ) + if not should_adjust: + yield + return + + weight_memory_before = torch.cuda.memory_allocated() + yield + target_weight_bytes = torch.cuda.memory_allocated() - weight_memory_before + model.mem_fraction = get_mtp_adjusted_mem_fraction( + mem_fraction=model.mem_fraction, + target_weight_bytes=target_weight_bytes, + target_layer_num=model.config["n_layer"], + mtp_layer_num=get_mtp_weight_layer_num(), + ) + + def load_config(weight_dir_): """ Load model configuration from the specified directory From d09c5c928d40dbe408183f14035357f23c4837d6 Mon Sep 17 00:00:00 2001 From: sufubao Date: Fri, 21 Aug 2026 18:12:22 +0800 Subject: [PATCH 081/103] fix(dspark): support partial rotary checkpoints --- lightllm/models/qwen3_5_dspark/model.py | 8 +- .../layer_infer/transformer_layer_infer.py | 3 + .../models/test_qwen3_dspark_model_output.py | 77 +++++++++++++++++-- 3 files changed, 80 insertions(+), 8 deletions(-) diff --git a/lightllm/models/qwen3_5_dspark/model.py b/lightllm/models/qwen3_5_dspark/model.py index aba188d310..865527597c 100644 --- a/lightllm/models/qwen3_5_dspark/model.py +++ b/lightllm/models/qwen3_5_dspark/model.py @@ -15,9 +15,11 @@ def _init_config(self): if "rope_theta" in rope_parameters and "rope_theta" not in self.config: self.config["rope_theta"] = rope_parameters["rope_theta"] - # Match the draft checkpoint's 1D full-head RoPE layout. - self.config["rope_scaling"] = None - self.config["partial_rotary_factor"] = 1.0 + # The draft is Qwen3-style and owns a 1D rotary cache. Released + # checkpoints rotate the full head, while target-shaped custom + # checkpoints can retain Qwen3.5's partial rotary layout. + self.config["rope_scaling"] = rope_parameters + self.config["partial_rotary_factor"] = rope_parameters.get("partial_rotary_factor", 1.0) def _init_custom(self): # Draft and target use different rotary shapes, so the draft owns its rotary cache. diff --git a/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py index 610ebb6641..84ee0b4059 100644 --- a/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py @@ -18,6 +18,7 @@ class Qwen3DFlashTransformerLayerInfer(LlamaTransformerLayerInfer): def __init__(self, layer_num, network_config): super().__init__(layer_num, network_config) self.head_dim_ = network_config["head_dim"] + self.partial_rotary_factor = network_config.get("partial_rotary_factor", 1.0) def context_forward( self, @@ -38,6 +39,7 @@ def context_forward( None, infer_state.position_cos, infer_state.position_sin, + partial_rotary_factor=self.partial_rotary_factor, ) self._post_cache_kv(cache_kv, infer_state, layer_weight) return input_embdings @@ -62,5 +64,6 @@ def _get_qkv(self, input, infer_state: Qwen3DFlashInferStateInfo, layer_weight: cache_kv[:, : self.tp_k_head_num_, :], infer_state.position_cos, infer_state.position_sin, + partial_rotary_factor=self.partial_rotary_factor, ) return q, cache_kv diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index 6e4a8f9929..6be526d8bb 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -6,6 +6,8 @@ from lightllm.common.basemodel import batch_objs from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel +from lightllm.models.qwen3_dflash.layer_infer import transformer_layer_infer as qwen3_dflash_layer_infer +from lightllm.models.qwen3_dflash.layer_infer.transformer_layer_infer import Qwen3DFlashTransformerLayerInfer from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel from lightllm.models.qwen3_dspark.layer_weights import pre_and_post_layer_weight as dspark_pre_post_weight from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer @@ -167,10 +169,9 @@ def init_dspark_config(self): "dflash_config": {"mask_token_id": 1}, "rope_parameters": { "rope_theta": 1_000_000, - "partial_rotary_factor": 0.25, - "mrope_interleaved": True, - "mrope_section": [11, 11, 10], - "rope_type": "default", + "factor": 32.0, + "original_max_position_embeddings": 8192, + "rope_type": "yarn", }, } @@ -179,12 +180,78 @@ def init_dspark_config(self): model._init_config() - assert model.config["rope_scaling"] is None + assert model.config["rope_scaling"] == model.config["rope_parameters"] assert model.config["rope_theta"] == 1_000_000 assert model.config["partial_rotary_factor"] == 1.0 assert model.config["mask_token_id"] == 1 +def test_qwen35_dspark_adapter_preserves_custom_partial_rotary_layout(monkeypatch): + def init_dspark_config(self): + self.config = { + "dflash_config": {"mask_token_id": 1}, + "rope_parameters": { + "mrope_interleaved": True, + "mrope_section": [11, 11, 10], + "partial_rotary_factor": 0.25, + "rope_theta": 10_000_000, + "rope_type": "default", + }, + } + + monkeypatch.setattr(Qwen3DSparkModel, "_init_config", init_dspark_config) + model = Qwen3_5DSparkModel.__new__(Qwen3_5DSparkModel) + + model._init_config() + + assert model.config["partial_rotary_factor"] == 0.25 + assert model.config["rope_scaling"] == model.config["rope_parameters"] + + +def test_qwen3_parallel_block_draft_applies_partial_rotary_factor(monkeypatch): + rotary_factors = [] + monkeypatch.setattr(qwen3_dflash_layer_infer, "qk_rmsnorm_forward", lambda *args, **kwargs: None) + monkeypatch.setattr( + qwen3_dflash_layer_infer, + "rotary_emb_fwd", + lambda *args, **kwargs: rotary_factors.append(kwargs["partial_rotary_factor"]), + ) + + class Projection: + def __init__(self, width): + self.width = width + + def mm(self, input, **kwargs): + return input.new_zeros((input.shape[0], self.width)) + + class QKNorm: + k_weight = object() + + def __call__(self, *args, **kwargs): + return None + + layer = Qwen3DFlashTransformerLayerInfer.__new__(Qwen3DFlashTransformerLayerInfer) + layer.tp_q_head_num_ = 1 + layer.tp_k_head_num_ = 1 + layer.tp_v_head_num_ = 1 + layer.head_dim_ = 8 + layer.eps_ = 1e-6 + layer.partial_rotary_factor = 0.25 + layer._post_cache_kv = lambda *args, **kwargs: None + infer_state = SimpleNamespace(position_cos=object(), position_sin=object()) + layer_weight = SimpleNamespace( + q_proj=Projection(8), + kv_proj=Projection(16), + qk_norm_weight_=QKNorm(), + ) + inputs = torch.zeros((2, 4)) + + layer.context_forward(inputs, infer_state, layer_weight) + layer._get_qkv(inputs, infer_state, layer_weight) + + assert rotary_factors == [0.25, 0.25] + + def test_vanilla_markov_local_sampling_matches_full_logits(): post_infer = Qwen3DSparkPostLayerInfer.__new__(Qwen3DSparkPostLayerInfer) post_infer.block_size_ = 3 From 29ec9f2f4c6612ce06c51229a0777ee0cc325afb Mon Sep 17 00:00:00 2001 From: root Date: Fri, 21 Aug 2026 10:42:22 +0000 Subject: [PATCH 082/103] dspark infer fix --- lightllm/models/qwen3_dspark/model.py | 30 ++++++++++++++++++++++++++- 1 file changed, 29 insertions(+), 1 deletion(-) diff --git a/lightllm/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py index c857248c98..aba99c9851 100644 --- a/lightllm/models/qwen3_dspark/model.py +++ b/lightllm/models/qwen3_dspark/model.py @@ -1,5 +1,10 @@ -from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel +from types import SimpleNamespace + +import torch + +from lightllm.common.basemodel.batch_objs import ModelOutput from lightllm.models.draft_registry import DraftModelRegistry +from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer from lightllm.models.qwen3_dspark.layer_weights.pre_and_post_layer_weight import Qwen3DSparkPreAndPostLayerWeight @@ -15,3 +20,26 @@ class Qwen3DSparkModel(Qwen3DFlashModel): pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight post_layer_infer_class = Qwen3DSparkPostLayerInfer + + def _decode(self, model_input): + if model_input.mtp_draft_input_hiddens is None: + return super()._decode(model_input) + + assert model_input.mtp_draft_input_hiddens.shape[0] == model_input.batch_size + + # This path only projects target hidden and writes its RoPE'd KV. It + # deliberately avoids the full decode infer state and attention setup. + position_ids = model_input.b_seq_len - 1 + infer_state = SimpleNamespace( + mtp_draft_input_hiddens=model_input.mtp_draft_input_hiddens, + position_cos=torch.index_select(self._cos_cached, 0, position_ids), + position_sin=torch.index_select(self._sin_cached, 0, position_ids), + mem_manager=self.mem_manager, + mem_index=model_input.mem_indexes, + ) + + hidden = self.pre_infer.context_forward(None, infer_state, self.pre_post_weight) + for layer, layer_weight in zip(self.layers_infer, self.trans_layers_weight): + hidden = layer.context_forward(hidden, infer_state, layer_weight) + + return ModelOutput(logits=hidden.new_empty((model_input.batch_size, 1))) From 481b344e58bf220b84e574a8feaf6bae4d687ce8 Mon Sep 17 00:00:00 2001 From: shihaobai <1798930569@qq.com> Date: Fri, 21 Aug 2026 18:55:12 +0800 Subject: [PATCH 083/103] feat: add qwen3.5 1p1d dspark launch script --- test/start_scripts/qwen35/qwen35_pd_1p1d.sh | 124 ++++++++++++++++++++ 1 file changed, 124 insertions(+) create mode 100755 test/start_scripts/qwen35/qwen35_pd_1p1d.sh diff --git a/test/start_scripts/qwen35/qwen35_pd_1p1d.sh b/test/start_scripts/qwen35/qwen35_pd_1p1d.sh new file mode 100755 index 0000000000..35e76780bf --- /dev/null +++ b/test/start_scripts/qwen35/qwen35_pd_1p1d.sh @@ -0,0 +1,124 @@ +#!/usr/bin/env bash +set -euo pipefail + +if [[ "$#" -lt 4 || "$#" -gt 5 ]]; then + echo "Usage: $0 [chat_template]" >&2 + exit 2 +fi + +PORT="$1" +MODEL_DIR="$2" +CPU_CACHE_SIZE="$3" +DSPARK_MODEL_DIR="$4" +CHAT_TEMPLATE_ARGS=() +if [[ -n "${5:-}" ]]; then + CHAT_TEMPLATE_ARGS=(--chat_template "$5") +fi + +# export PYTORCH_ALLOC_CONF=expandable_segments:True +export LOADWORKER=8 +export LIGHTLLM_TRITON_AUTOTUNE_LEVEL=1 +export LIGHTLLM_ANTHROPIC_ENABLE_PDF_PARSING=1 +export LIGHTLLM_LOG_LEVEL=debug + +# Keep same-host PD WebSocket traffic away from HTTP/WebSocket proxies. +export NO_PROXY="${NO_PROXY:+${NO_PROXY},}127.0.0.1,localhost" +export no_proxy="${NO_PROXY}" + +P_COMMON_ARGS=( + --model_dir "${MODEL_DIR}" + --model_name qwen35_27b + --graph_max_batch_size 8 + --running_max_req_size 8 + --mem_fraction 0.80 + --max_image_token_count 4096 + --max_image_pixels 3686400 + --batch_max_tokens 8192 + --linear_att_cache_size 500 + --linear_att_hash_page_size 2048 + --linear_att_page_block_num 8 + --quant_type fp8w8a8-pt-sgl + --mtp_mode dspark + --mtp_draft_model_dir "${DSPARK_MODEL_DIR}" + --mtp_step 1 + "${CHAT_TEMPLATE_ARGS[@]}" + --pd_trans_mode nccl + --pd_kv_page_size 4096 + --pd_master_ip 127.0.0.1 + --pd_master_port "${PORT}" + --enable_prefill_cudagraph +) + +D_COMMON_ARGS=( + --model_dir "${MODEL_DIR}" + --model_name qwen35_27b + --graph_max_batch_size 64 + --running_max_req_size 64 + --mem_fraction 0.80 + --max_image_token_count 4096 + --max_image_pixels 3686400 + --batch_max_tokens 512 + --linear_att_cache_size 500 + --quant_type fp8w8a8-pt-sgl + --mtp_mode dspark + --mtp_draft_model_dir "${DSPARK_MODEL_DIR}" + --mtp_step 5 + "${CHAT_TEMPLATE_ARGS[@]}" + --pd_trans_mode nccl + --pd_kv_page_size 4096 + --pd_master_ip 127.0.0.1 + --pd_master_port "${PORT}" +) + +PIDS=() + +cleanup() { + kill -TERM "${PIDS[@]}" 2>/dev/null || true + wait "${PIDS[@]}" 2>/dev/null || true +} + +trap cleanup EXIT +trap 'exit 130' INT +trap 'exit 143' TERM + +# Prefill: TP4 on GPUs 0-3. +CUDA_VISIBLE_DEVICES=0,1,2,3 python -m lightllm.server.api_server \ + "${P_COMMON_ARGS[@]}" \ + --run_mode prefill \ + --enable_cpu_cache \ + --cpu_cache_storage_size "${CPU_CACHE_SIZE}" \ + --tp 4 \ + --visual_dp 4 \ + --host 0.0.0.0 \ + --port 28761 \ + --nccl_port 29761 & +PIDS+=("$!") + +# Decode: TP4 on GPUs 4-7. +CUDA_VISIBLE_DEVICES=4,5,6,7 python -m lightllm.server.api_server \ + "${D_COMMON_ARGS[@]}" \ + --run_mode decode \ + --tp 4 \ + --host 0.0.0.0 \ + --port 28762 \ + --nccl_port 29762 & +PIDS+=("$!") + +# Start PD Master last, after Prefill and Decode have started. +python -m lightllm.server.api_server \ + --model_dir "${MODEL_DIR}" \ + --model_name qwen35_27b \ + --run_mode pd_master \ + --pd_master_mode 1p1d \ + --host 0.0.0.0 \ + --port "${PORT}" \ + --max_image_token_count 4096 \ + --max_image_pixels 3686400 \ + "${CHAT_TEMPLATE_ARGS[@]}" & +PIDS+=("$!") + +echo "1P1D D-Spark service is starting at http://127.0.0.1:${PORT}" + +# A fixed 1P1D deployment is incomplete as soon as any component exits. +wait -n "${PIDS[@]}" +exit 1 From 9f2fe17e0fbf77c308aedec87e655df22e246e60 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 21 Aug 2026 11:47:02 +0000 Subject: [PATCH 084/103] refactor: reuse common speculative engine for dp --- .../triton_kernel/select_mtp_rows.py | 4 +- .../mode_backend/dp_backend/impl.py | 263 ++++++++---------- .../model_infer/mtp_speculative/__init__.py | 3 +- .../model_infer/mtp_speculative/dp_engine.py | 58 ---- .../dp_overlap_proposers/base.py | 14 +- .../dp_overlap_proposers/eagle3.py | 30 -- .../dp_overlap_proposers/eagle_no_att.py | 30 -- .../dp_overlap_proposers/eagle_utils.py | 51 +++- .../dp_overlap_proposers/eagle_with_att.py | 30 -- .../dp_overlap_proposers/vanilla_no_att.py | 61 ---- .../dp_overlap_proposers/vanilla_with_att.py | 78 ------ .../mtp_speculative/dp_planner/__init__.py | 14 - .../mtp_speculative/dp_planner/base.py | 11 - .../mtp_speculative/dp_planner/fixed.py | 11 - .../mtp_speculative/dp_proposers/__init__.py | 39 --- .../mtp_speculative/dp_proposers/base.py | 7 - .../mtp_speculative/dp_proposers/eagle3.py | 44 --- .../dp_proposers/eagle_no_att.py | 41 --- .../dp_proposers/eagle_utils.py | 199 ------------- .../dp_proposers/eagle_with_att.py | 41 --- .../dp_proposers/vanilla_no_att.py | 72 ----- .../dp_proposers/vanilla_with_att.py | 156 ----------- .../model_infer/mtp_speculative/engine.py | 24 +- .../mtp_speculative/planner/base.py | 3 +- .../test_dp_overlap_spec_engine.py | 110 ++++---- .../mtp_speculative/test_planner.py | 62 +---- .../mtp_speculative/test_vanilla_no_att.py | 67 ++++- 27 files changed, 315 insertions(+), 1208 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_engine.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_planner/base.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_planner/fixed.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/base.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_utils.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py diff --git a/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py b/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py index 1ed2877f12..5e230ce855 100644 --- a/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py +++ b/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py @@ -111,7 +111,6 @@ def select_accepted_tail_rows( """Select one accepted-tail row per request in a single CUDA kernel.""" req_num = b_req_mtp_start_loc.shape[0] - assert req_num > 0 assert input_ids.is_cuda assert hidden.ndim == 2 and hidden.shape[0] == input_ids.shape[0] assert hidden.shape[1] > 0 @@ -140,6 +139,9 @@ def select_accepted_tail_rows( b_shared_radix_node_id=b_shared_radix_node_id.new_empty((req_num,)), b_position_delta=b_position_delta.new_empty((req_num,)), ) + if req_num == 0: + return selected + block_hidden = 1024 pipeline_stages = 3 grid = (req_num,) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 9c99eed04b..6d34dba034 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -1,6 +1,7 @@ import torch import time from typing import List, Tuple +from lightllm.common.basemodel.triton_kernel.mtp_utils import gen_b_req_mtp_start_loc from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.common.basemodel.batch_objs import ModelOutput, ModelInput from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq @@ -15,7 +16,7 @@ from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.envs_utils import get_env_start_args from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager -from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree @@ -52,9 +53,6 @@ def __init__(self) -> None: ) else: self.decode = self.decode_mtp - self._draft_decode_func = ( - self._draft_decode_eagle if self.uses_autoregressive_drafter else self._draft_decode_vanilla - ) else: if self.enable_prefill_microbatch_overlap: self.prefill = self.prefill_overlap @@ -75,7 +73,15 @@ def init_spec_engine(self): spec_mode=self.args.mtp_mode, enable_dynmaic_mtp=self.args.mtp_dynamic_verify, ) - self.spec_engine = DPSpecEngine(**engine_kwargs) + # 非 overlap DP 与普通后端复用同一个 SpecEngine。固定 draft 深度 + # 保证所有 DP rank 执行相同数量的 draft forward;空 rank 则由 + # allow_empty_batch 进入完整的 dummy forward 流程。 + self.spec_engine = SpecEngine( + backend=self, + spec_mode=self.args.mtp_mode, + enable_dynmaic_mtp=False, + allow_empty_batch=True, + ) self.dp_overlap_spec_engine = DPOverlapSpecEngine(**engine_kwargs) self.prefill_draft_engine = ( self.dp_overlap_spec_engine if self.enable_prefill_microbatch_overlap else self.spec_engine @@ -472,17 +478,12 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] mask_func=None, ) - # mtp kv fill - target_next_token_ids_gpu = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input.batch_size, - copy_len=req_num, - device=model_input.b_req_idx.device, - ) - self.prefill_draft_engine.fill_draft_model_kv_state( + # BaseModel 已负责空 batch 的内部 padding,这里直接把真实 target + # 输出交给与非 DP 路径相同的 SpecEngine。 + self.spec_engine.fill_draft_model_kv_state( target_model_input=model_input, target_model_output=model_output, - target_next_token_ids=target_next_token_ids_gpu, + target_next_token_ids=next_token_ids, ) if req_num > 0: g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=b_req_idx, reqs=run_reqs) @@ -519,42 +520,59 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] return def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): + """复用普通 SpecEngine 执行 DP speculative draft-and-verify。""" + model_input, run_reqs = prepare_decode_inputs(req_objs=decode_reqs) - b_mtp_index_cpu = model_input.b_mtp_index - req_num = len(run_reqs) + spec_engine = self.spec_engine + req_num = len(decode_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): + spec_plan = spec_engine.plan_decode( + model_input=model_input, + decode_reqs=decode_reqs, + ) + model_input, async_selected_row_mask_cpu = spec_engine.prepare_decode_model_input( + model_input=model_input, + req_num=req_num, + plan=spec_plan, + ) model_output = self.model.forward(model_input) - mtp_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None - if req_num > 0: - logits = model_output.logits[0:req_num, :] - b_mtp_index_cpu = b_mtp_index_cpu[0:req_num] - b_req_idx = model_input.b_req_idx[0:req_num] - next_token_ids, next_token_logprobs = sample(logits, run_reqs, self.eos_id) - next_token_ranks = self._get_next_token_ranks(logits, next_token_ids) - ( - next_token_ids_cpu, - next_token_logprobs_cpu, - next_token_ranks_cpu, - ) = self._async_copy_next_token_infos_to_pin_mem(next_token_ids, next_token_logprobs, next_token_ranks) + if async_selected_row_mask_cpu is not None: + async_selected_row_mask_cpu.wait() + selected_rows = async_selected_row_mask_cpu.tensor.tolist() + run_reqs = [req for req, selected in zip(run_reqs, selected_rows) if selected] - # verify the next_token_ids - b_req_mtp_start_loc = [ - index for index, mtp_index in enumerate(b_mtp_index_cpu.tolist()) if mtp_index == 0 - ] - b_req_mtp_start_loc = g_pin_mem_manager.gen_from_list( - key="b_req_mtp_start_loc", - data=b_req_mtp_start_loc, - dtype=torch.int32, - ).cuda(non_blocking=True) + next_token_ids = torch.empty( + (0,), + dtype=torch.int64, + device=model_input.b_req_idx.device, + ) + b_req_mtp_start_loc = torch.empty( + (0,), + dtype=torch.int32, + device=model_input.b_req_idx.device, + ) + mtp_accept_len = torch.empty_like(b_req_mtp_start_loc) + if req_num > 0: + next_token_ids, next_token_logprobs = sample( + model_output.logits, + run_reqs, + self.eos_id, + ) + next_token_ranks = self._get_next_token_ranks(model_output.logits, next_token_ids) + + b_req_mtp_start_loc = gen_b_req_mtp_start_loc( + b_mtp_index=model_input.b_mtp_index, + num_reqs=req_num, + ) mtp_accept_len, accepted_index = mtp_utils.verify_mtp_tokens( backend=self, next_token_ids=next_token_ids, - b_req_idx=b_req_idx, + b_req_idx=model_input.b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, - b_mtp_index=model_input.b_mtp_index[0:req_num], + b_mtp_index=model_input.b_mtp_index, ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", @@ -564,23 +582,40 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): key="mtp_accept_len", gpu_tensor=mtp_accept_len, ) + g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( + b_req_idx=model_input.b_req_idx, + next_token_ids=next_token_ids, + mask=accepted_index == 1, + ) verify_event = torch.cuda.Event() verify_event.record() - extra_mem_indexes_cpu = self._draft_decode_func( - model_input=model_input, - model_output=model_output, - next_token_ids=next_token_ids, + proposal = spec_engine.propose_next( + target_model_input=model_input, + target_model_output=model_output, + target_next_token_ids=next_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, - mtp_accept_len=mtp_accept_len, - req_num=req_num, + draft_step=spec_plan.draft_step, + accept_len=mtp_accept_len, ) if req_num > 0: - g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( - b_req_idx=b_req_idx, + mtp_utils.scatter_mtp_next_tokens( + backend=self, + proposal=proposal, + target_next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_req_idx=model_input.b_req_idx, + mtp_accept_len=mtp_accept_len, + ) + ( + next_token_ids_cpu, + next_token_logprobs_cpu, + next_token_ranks_cpu, + ) = self._async_copy_next_token_infos_to_pin_mem( next_token_ids=next_token_ids, - mask=accepted_index == 1, + next_token_logprobs=next_token_logprobs, + next_token_ranks=next_token_ranks, ) sync_event = torch.cuda.Event() @@ -589,22 +624,34 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): if req_num > 0: # 第二阶段 event_pack.notify_post_handle_and_wait_pre_post_handle() - verify_event.synchronize() + if spec_plan.skip_verify_sync: + verify_ok_reqs = run_reqs + else: + verify_event.synchronize() + verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] + + update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) + + # 第三阶段 + event_pack.notify_forward_and_wait_post_handle() + sync_event.synchronize() + + spec_engine.update_planner_statics( + plan=spec_plan, + proposal=proposal, + req_num=req_num, + accept_lengths_cpu=mtp_accept_len_cpu, + ) mtp_utils.record_request_mtp_metrics( backend=self, decode_reqs=decode_reqs, accept_lengths_cpu=mtp_accept_len_cpu, verify_run_reqs=run_reqs, ) - verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] - update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) - # 第三阶段 - event_pack.notify_forward_and_wait_post_handle() - sync_event.synchronize() - extra_mem_indexes_cpu.append( + proposal.extra_mem_indexes_cpu.append( MtpMemIndexesToFree( - mem_indexes_cpu=model_input.mem_indexes_cpu[0:req_num], + mem_indexes_cpu=model_input.mem_indexes_cpu, free_mask_cpu=accepted_index_cpu == 0, ), ) @@ -620,7 +667,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): ) mtp_utils.free_mem_indexes( backend=self, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, ) # 第四阶段 @@ -628,105 +675,13 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): else: event_pack.notify_post_handle_and_wait_pre_post_handle() event_pack.notify_forward_and_wait_post_handle() - event_pack.notify_pre_post_handle() - return - - def _draft_decode_vanilla( - self, - model_input: ModelInput, - model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - mtp_accept_len: torch.Tensor, - req_num: int, - ): - if b_req_mtp_start_loc is None: - b_req_mtp_start_loc = torch.empty((0,), dtype=torch.int32, device=model_input.b_req_idx.device) - mtp_accept_len = torch.empty_like(b_req_mtp_start_loc) - padded_next_token_ids = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input.batch_size, - copy_len=req_num, - device=model_input.b_req_idx.device, - ) - proposal = self.decode_draft_engine.propose_next( - target_model_input=model_input, - target_model_output=model_output, - target_next_token_ids=padded_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=mtp_accept_len, - ) - - if req_num > 0: - mtp_utils.scatter_mtp_next_tokens( - backend=self, - proposal=proposal, - target_next_token_ids=next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=model_input.b_req_idx[:req_num], - mtp_accept_len=mtp_accept_len, - ) - return proposal.extra_mem_indexes_cpu - - def _draft_decode_eagle( - self, - model_input: ModelInput, - model_output: ModelOutput, - next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - mtp_accept_len: torch.Tensor, - req_num: int, - ): - verify_width = self.max_draft_step + 1 - real_request_num = req_num // verify_width - request_capacity = model_input.batch_size // verify_width - - padded_next_token_ids = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input.batch_size, - copy_len=req_num, - device=model_input.b_req_idx.device, - ) - padded_start_locs = torch.arange( - 0, - model_input.batch_size, - verify_width, - dtype=torch.int32, - device=model_input.b_req_idx.device, - ) - padded_accept_len = torch.ones( - (request_capacity,), - dtype=torch.int32, - device=model_input.b_req_idx.device, - ) - if real_request_num > 0: - padded_accept_len[:real_request_num].copy_(mtp_accept_len) - - # DP keeps the target verify layout padded for collective shape - # agreement. The proposer still follows the common topology: one - # full-row extend, followed by autoregressive drafting over one row per - # (real or HOLD) request. - proposal = self.decode_draft_engine.propose_next( - target_model_input=model_input, - target_model_output=model_output, - target_next_token_ids=padded_next_token_ids, - b_req_mtp_start_loc=padded_start_locs, - accept_len=padded_accept_len, - ) - proposal.token_ids = proposal.token_ids[:real_request_num] - if getattr(proposal, "schedule_scores", None) is not None: - proposal.schedule_scores = proposal.schedule_scores[:real_request_num] - - if req_num > 0: - mtp_utils.scatter_mtp_next_tokens( + sync_event.synchronize() + mtp_utils.free_mem_indexes( backend=self, - proposal=proposal, - target_next_token_ids=next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=model_input.b_req_idx[:req_num], - mtp_accept_len=mtp_accept_len, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, ) - return proposal.extra_mem_indexes_cpu + event_pack.notify_pre_post_handle() + return def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): ( diff --git a/lightllm/server/router/model_infer/mtp_speculative/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/__init__.py index 4944600fbd..66dc8ab2fc 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/__init__.py +++ b/lightllm/server/router/model_infer/mtp_speculative/__init__.py @@ -1,6 +1,5 @@ -from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine -__all__ = ["DPSpecEngine", "DPOverlapSpecEngine", "SpecEngine"] +__all__ = ["DPOverlapSpecEngine", "SpecEngine"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py deleted file mode 100644 index 6b9a109c3f..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_engine.py +++ /dev/null @@ -1,58 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, Optional - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_planner import BaseDpPlanner, build_dp_planner -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers import build_dp_spec_proposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal - -if TYPE_CHECKING: - from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend - - -class DPSpecEngine: - """普通 DP prefill/decode 使用的 MTP engine。""" - - def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: - self.proposer: BaseDpProposer = build_dp_spec_proposer( - spec_mode=spec_mode, - backend=backend, - enable_dynmaic_mtp=enable_dynmaic_mtp, - ) - self.planner: BaseDpPlanner = build_dp_planner(backend=backend) - - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - self.proposer.fill_draft_model_kv_state( - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - ) - - def propose_next( - self, - target_model_input: ModelInput, # batch_size = verify_batch_size - target_model_output: ModelOutput, # logits: [verify_batch_size, vocab_size] - target_next_token_ids: torch.Tensor, # [verify_batch_size] - b_req_mtp_start_loc: torch.Tensor, # [req_num] - accept_len: Optional[torch.Tensor] = None, # [req_num] - ) -> SpecProposal: - return self.proposer.propose_next( - target_model_input=target_model_input, - target_model_output=target_model_output, - target_next_token_ids=target_next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - draft_step=self.planner.get_draft_step(), - accept_len=accept_len, - ) - - -__all__ = ["DPSpecEngine"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py index b085dc85bd..7e3877894c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -1,13 +1,21 @@ from abc import ABC, abstractmethod +from typing import TYPE_CHECKING import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend -class BaseDpOverlapProposer(BaseSpecProposer, ABC): - """DP proposer 的完整接口,扩展双 microbatch overlap 操作。""" + +class BaseDpOverlapProposer(ABC): + """双 microbatch DP-overlap proposer 的独立接口。""" + + def __init__(self, *, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> None: + self.backend = backend + self.enable_dynmaic_mtp = bool(enable_dynmaic_mtp) @abstractmethod def fill_draft_model_kv_state_overlap( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index 9e947893ab..69fc25928e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -4,9 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, - fill_eagle_draft_model_kv_state, propose_next_dp_eagle_autoregressive_overlap, - propose_next_eagle, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal @@ -17,14 +15,6 @@ class DpOverlapEagle3Proposer(BaseDpOverlapProposer): def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) - def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, @@ -44,26 +34,6 @@ def fill_draft_model_kv_state_overlap( target_next_token_ids1, ) - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> EagleSpecProposal: - return propose_next_eagle( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, - self._map_draft_token_ids, - ) - def propose_next_overlap( self, target_model_input0: ModelInput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 8383f71ec1..fe151f6a83 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -4,9 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, - fill_eagle_draft_model_kv_state, propose_next_dp_eagle_fixed_layout_overlap, - propose_next_eagle, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal @@ -14,14 +12,6 @@ class DpOverlapEagleNoAttProposer(BaseDpOverlapProposer): """DP ``eagle_no_att`` proposer。""" - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) - def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, @@ -41,26 +31,6 @@ def fill_draft_model_kv_state_overlap( target_next_token_ids1, ) - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> EagleSpecProposal: - return propose_next_eagle( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, - lambda token_ids: token_ids, - ) - def propose_next_overlap( self, target_model_input0: ModelInput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 27af2c5dfb..72f3fcac08 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -7,17 +7,54 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_utils import ( - _prepare_eagle_prefill_inputs, - fill_eagle_draft_model_kv_state, - generate_eagle_token_ids, - prepare_eagle_verify_decode_input, - propose_next_eagle, -) from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + + +def _prepare_eagle_prefill_inputs( + model_input: ModelInput, + b_next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, +) -> None: + """构造 DP-overlap EAGLE draft model 的 prefill 输入。""" + + model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( + key="dp_overlap_eagle_prefill_b_is_decode_req", + shape=model_input.b_req_idx.shape, + fill_value=False, + dtype=torch.bool, + ) + model_input.input_ids = gen_mtp_new_input_ids( + input_ids=model_input.input_ids, + b_next_token_ids=b_next_token_ids, + b_seq_len=model_input.b_seq_len, + b_ready_cache_len=model_input.b_ready_cache_len, + ) + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens + + +def generate_eagle_token_ids( + proposer: BaseDpOverlapProposer, + model_output: ModelOutput, + map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], +) -> torch.Tensor: + return map_draft_token_ids(proposer.backend._gen_argmax_token_ids(model_output)) + + +def prepare_eagle_verify_decode_input( + model_input: ModelInput, + input_ids: torch.Tensor, + target_hidden: torch.Tensor, +) -> None: + """复用 target verify 布局构造 DP-overlap drafter 输入。""" + + assert not model_input.is_prefill + model_input.input_ids = input_ids + model_input.mtp_draft_input_hiddens = target_hidden def fill_dp_eagle_draft_model_kv_state_overlap( diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index efb179c11a..c470a72445 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -4,9 +4,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, - fill_eagle_draft_model_kv_state, propose_next_dp_eagle_fixed_layout_overlap, - propose_next_eagle, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal @@ -14,14 +12,6 @@ class DpOverlapEagleWithAttProposer(BaseDpOverlapProposer): """DP ``eagle_with_att`` proposer。""" - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) - def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, @@ -41,26 +31,6 @@ def fill_draft_model_kv_state_overlap( target_next_token_ids1, ) - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> EagleSpecProposal: - return propose_next_eagle( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, - lambda token_ids: token_ids, - ) - def propose_next_overlap( self, target_model_input0: ModelInput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 2489dd9d26..e29f805430 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -10,14 +10,6 @@ class DpOverlapVanillaNoAttProposer(BaseDpOverlapProposer): """DP ``vanilla_no_att`` proposer。""" - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - pass - def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, @@ -29,59 +21,6 @@ def fill_draft_model_kv_state_overlap( ) -> None: pass - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> VanillaSpecProposal: - req_num = int(b_req_mtp_start_loc.shape[0]) - proposal_token_ids_by_step = [] - schedule_scores_by_step = [] - - if draft_step == 0: - return VanillaSpecProposal( - token_ids=target_next_token_ids.new_empty((req_num, 0)), - extra_mem_indexes_cpu=[], - schedule_scores=( - torch.empty((req_num, 0), dtype=torch.float32, device=target_next_token_ids.device) - if self.enable_dynmaic_mtp - else None - ), - ) - - assert accept_len is not None - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_token_ids = target_next_token_ids - draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_input = copy.copy(target_model_input) - - for step in range(draft_step): - draft_input.input_ids = draft_token_ids - draft_input.mtp_draft_input_hiddens = draft_hidden - draft_output = self.backend.draft_models[step].forward(draft_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - - if self.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) - selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) - schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) - else: - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) - proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) - - proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) - schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None - return VanillaSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=schedule_scores, - ) - def propose_next_overlap( self, target_model_input0: ModelInput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 8bbfcd7214..c1f8e3d6de 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -13,29 +13,6 @@ class DpOverlapVanillaWithAttProposer(BaseDpOverlapProposer): """DP ``vanilla_with_att`` proposer。""" - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - assert target_model_input.is_prefill - assert target_model_input.b_position_delta is None - assert target_next_token_ids.shape == target_model_input.b_req_idx.shape - - draft_input = copy.copy(target_model_input) - draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_token_ids = target_next_token_ids - for draft_model in self.backend.draft_models: - draft_input = self._prepare_mtp_prefill_inputs( - model_input=draft_input, - b_next_token_ids=draft_token_ids, - mtp_draft_input_hiddens=draft_hidden, - ) - draft_output = draft_model.forward(draft_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - def fill_draft_model_kv_state_overlap( self, target_model_input0: ModelInput, @@ -71,61 +48,6 @@ def fill_draft_model_kv_state_overlap( draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> VanillaSpecProposal: - req_num = int(b_req_mtp_start_loc.shape[0]) - proposal_token_ids_by_step = [] - schedule_scores_by_step = [] - - assert not target_model_input.is_prefill - assert accept_len is not None - assert accept_len.shape == (req_num,) - assert draft_step == self.backend.max_draft_step - assert len(self.backend.draft_models) == draft_step - - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_token_ids = target_next_token_ids - draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_input = copy.copy(target_model_input) - - for step in range(draft_step): - draft_input.input_ids = draft_token_ids - draft_input.mtp_draft_input_hiddens = draft_hidden - draft_output = self.backend.draft_models[step].forward(draft_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - - if self.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) - selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) - schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) - else: - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) - proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) - - if step + 1 < draft_step: - draft_token_ids = overlay_chained_mtp_decode_input( - input_ids=draft_input.input_ids, - draft_token_ids=draft_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=accept_len, - ) - - proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) - schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None - return VanillaSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=schedule_scores, - ) - def propose_next_overlap( self, target_model_input0: ModelInput, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py deleted file mode 100644 index d57a0eecd1..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/__init__.py +++ /dev/null @@ -1,14 +0,0 @@ -from typing import TYPE_CHECKING - -from lightllm.server.router.model_infer.mtp_speculative.dp_planner.base import BaseDpPlanner -from lightllm.server.router.model_infer.mtp_speculative.dp_planner.fixed import FixedDpPlanner - -if TYPE_CHECKING: - from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend - - -def build_dp_planner(*, backend: "ModeBackend") -> BaseDpPlanner: - return FixedDpPlanner(draft_step=backend.max_draft_step) - - -__all__ = ["BaseDpPlanner", "FixedDpPlanner", "build_dp_planner"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/base.py deleted file mode 100644 index bd20e39310..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/base.py +++ /dev/null @@ -1,11 +0,0 @@ -from abc import ABC, abstractmethod - - -class BaseDpPlanner(ABC): - """普通 DP draft 配置的基础规划接口。""" - - @abstractmethod - def get_draft_step(self) -> int: - """返回当前 DP decode proposal 使用的 draft step。""" - - raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/dp_planner/fixed.py deleted file mode 100644 index 3fdf4a7b63..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_planner/fixed.py +++ /dev/null @@ -1,11 +0,0 @@ -from lightllm.server.router.model_infer.mtp_speculative.dp_planner.base import BaseDpPlanner - - -class FixedDpPlanner(BaseDpPlanner): - """为 DP 固定 shape decode 返回启动时配置的 draft step。""" - - def __init__(self, draft_step: int) -> None: - self.draft_step = int(draft_step) - - def get_draft_step(self) -> int: - return self.draft_step diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py deleted file mode 100644 index 1155b67d81..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/__init__.py +++ /dev/null @@ -1,39 +0,0 @@ -from typing import TYPE_CHECKING - -if TYPE_CHECKING: - from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend - from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer - - -def build_dp_spec_proposer(*, spec_mode: str, backend: "ModeBackend", enable_dynmaic_mtp: bool) -> "BaseDpProposer": - if spec_mode == "vanilla_with_att": - from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_with_att import ( - DpVanillaWithAttProposer, - ) - - return DpVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) - if spec_mode == "vanilla_no_att": - from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_no_att import ( - DpVanillaNoAttProposer, - ) - - return DpVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) - if spec_mode == "eagle_with_att": - from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_with_att import ( - DpEagleWithAttProposer, - ) - - return DpEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) - if spec_mode == "eagle_no_att": - from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_no_att import DpEagleNoAttProposer - - return DpEagleNoAttProposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) - if spec_mode == "eagle3": - from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle3 import DpEagle3Proposer - - return DpEagle3Proposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) - - raise ValueError(f"unsupported DP speculative mode: {spec_mode}") - - -__all__ = ["build_dp_spec_proposer"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/base.py deleted file mode 100644 index aada294da8..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/base.py +++ /dev/null @@ -1,7 +0,0 @@ -from abc import ABC - -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer - - -class BaseDpProposer(BaseSpecProposer, ABC): - """普通 DP prefill/decode proposer 的基础接口。""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py deleted file mode 100644 index 77e2a8e354..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle3.py +++ /dev/null @@ -1,44 +0,0 @@ -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_utils import ( - fill_eagle_draft_model_kv_state, - propose_next_eagle, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal - - -class DpEagle3Proposer(BaseDpProposer): - """普通 DP ``eagle3`` proposer。""" - - def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: - return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) - - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) - - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> EagleSpecProposal: - return propose_next_eagle( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, - self._map_draft_token_ids, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py deleted file mode 100644 index 16479336d9..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_no_att.py +++ /dev/null @@ -1,41 +0,0 @@ -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_utils import ( - fill_eagle_draft_model_kv_state, - propose_next_eagle, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal - - -class DpEagleNoAttProposer(BaseDpProposer): - """普通 DP ``eagle_no_att`` proposer。""" - - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) - - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> EagleSpecProposal: - return propose_next_eagle( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, - lambda token_ids: token_ids, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_utils.py deleted file mode 100644 index c647137904..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_utils.py +++ /dev/null @@ -1,199 +0,0 @@ -"""普通 DP 与 DP-overlap EAGLE proposer 共用的辅助函数。""" - -from __future__ import annotations - -import copy -from typing import Callable - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( - BaseSpecProposer, - MtpMemIndexesToFree, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager - - -def fill_eagle_draft_model_kv_state( - proposer: BaseSpecProposer, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, -) -> None: - """使用 target prefill 输出初始化 EAGLE draft state。""" - - _prepare_eagle_prefill_inputs( - model_input=target_model_input, - b_next_token_ids=target_next_token_ids, - mtp_draft_input_hiddens=target_model_output.mtp_collector.spec_hidden, - ) - proposer.backend.draft_models[0].forward(target_model_input) - - -def _prepare_eagle_prefill_inputs( - model_input: ModelInput, - b_next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: torch.Tensor, -) -> ModelInput: - model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( - key="eagle_mtp_prefill_b_is_decode_req", - shape=model_input.b_req_idx.shape, - fill_value=False, - dtype=torch.bool, - ) - model_input.input_ids = gen_mtp_new_input_ids( - input_ids=model_input.input_ids, - b_next_token_ids=b_next_token_ids, - b_seq_len=model_input.b_seq_len, - b_ready_cache_len=model_input.b_ready_cache_len, - ) - model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens - return model_input - - -def generate_eagle_token_ids( - proposer: BaseSpecProposer, - model_output: ModelOutput, - map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> torch.Tensor: - return map_draft_token_ids(proposer.backend._gen_argmax_token_ids(model_output)) - - -def generate_eagle_token_ids_and_prob( - proposer: BaseSpecProposer, - model_output: ModelOutput, - map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -): - draft_token_ids, draft_token_probs = proposer.backend._gen_argmax_token_ids_and_prob(model_output) - return map_draft_token_ids(draft_token_ids), draft_token_probs - - -def prepare_eagle_verify_decode_input( - model_input: ModelInput, - input_ids: torch.Tensor, - target_hidden: torch.Tensor, -) -> None: - """复用 target MTP decode 布局,为 drafter 的 verification KV commit 准备输入。""" - - assert not model_input.is_prefill - model_input.input_ids = input_ids - model_input.mtp_draft_input_hiddens = target_hidden - - -def propose_next_eagle( - proposer: BaseSpecProposer, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None, - map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> EagleSpecProposal: - """运行 EAGLE extend 后接单 token decode 的通用 proposal 流程。""" - - request_count = int(b_req_mtp_start_loc.shape[0]) - proposal_token_ids = target_next_token_ids.new_empty((request_count, draft_step)) - collect_schedule_scores = proposer.enable_dynmaic_mtp - schedule_scores = ( - torch.zeros( - (request_count, draft_step), - dtype=torch.float32, - device=target_next_token_ids.device, - ) - if collect_schedule_scores - else None - ) - if draft_step == 0: - return EagleSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=schedule_scores, - ) - - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_model = proposer.backend.draft_models[0] - position_delta = target_model_input.b_position_delta - assert position_delta is not None - prepare_eagle_verify_decode_input( - model_input=target_model_input, - input_ids=target_next_token_ids, - target_hidden=target_model_output.mtp_collector.spec_hidden, - ) - extend_output = draft_model.forward(target_model_input) - - accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) - if collect_schedule_scores: - draft_token_ids, draft_token_probs = generate_eagle_token_ids_and_prob( - proposer=proposer, - model_output=accepted_tail_output, - map_draft_token_ids=map_draft_token_ids, - ) - schedule_scores[:, 0] = draft_token_probs.float() - else: - draft_token_ids = generate_eagle_token_ids( - proposer=proposer, - model_output=accepted_tail_output, - map_draft_token_ids=map_draft_token_ids, - ) - proposal_token_ids[:, 0] = draft_token_ids - draft_hidden = extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows) - - if draft_step == 1: - return EagleSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=schedule_scores, - ) - - extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(request_count * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids.device, non_blocking=True) - draft_seq_lens = target_model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1 - max_kv_seq_len = target_model_input.max_kv_seq_len - draft_input = copy.copy(target_model_input) - draft_input.is_prefill = False - draft_input.batch_size = request_count - draft_input.b_req_idx = target_model_input.b_req_idx.index_select(0, accepted_tail_rows) - draft_input.b_mtp_index = torch.zeros_like(draft_input.b_req_idx) - draft_input.b_seq_len = draft_seq_lens - draft_input.b_position_delta = position_delta.index_select(0, accepted_tail_rows) - draft_input.b_shared_seq_len = target_model_input.b_shared_seq_len.index_select(0, accepted_tail_rows) - draft_input.b_shared_radix_node_id = target_model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) - if len(draft_input.multimodal_params) != request_count: - empty_multimodal_params = {"images": [], "audios": []} - draft_input.multimodal_params = [empty_multimodal_params] * request_count - - for step in range(1, draft_step): - mem_start = (step - 1) * request_count - draft_input.input_ids = draft_token_ids - draft_input.mtp_draft_input_hiddens = draft_hidden - draft_input.mem_indexes = extra_mem_indexes[mem_start : mem_start + request_count] - draft_input.max_kv_seq_len = max_kv_seq_len + step - draft_input.total_token_num = request_count * draft_input.max_kv_seq_len - draft_output = draft_model.forward(draft_input) - if collect_schedule_scores: - draft_token_ids, draft_token_probs = generate_eagle_token_ids_and_prob( - proposer=proposer, - model_output=draft_output, - map_draft_token_ids=map_draft_token_ids, - ) - schedule_scores[:, step] = draft_token_probs.float() - else: - draft_token_ids = generate_eagle_token_ids( - proposer=proposer, - model_output=draft_output, - map_draft_token_ids=map_draft_token_ids, - ) - proposal_token_ids[:, step] = draft_token_ids - draft_hidden = draft_output.mtp_collector.spec_hidden - draft_seq_lens.add_(1) - - return EagleSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], - schedule_scores=schedule_scores, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py deleted file mode 100644 index fa36357eda..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/eagle_with_att.py +++ /dev/null @@ -1,41 +0,0 @@ -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_utils import ( - fill_eagle_draft_model_kv_state, - propose_next_eagle, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal - - -class DpEagleWithAttProposer(BaseDpProposer): - """普通 DP ``eagle_with_att`` proposer。""" - - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - fill_eagle_draft_model_kv_state(self, target_model_input, target_model_output, target_next_token_ids) - - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> EagleSpecProposal: - return propose_next_eagle( - self, - target_model_input, - target_model_output, - target_next_token_ids, - b_req_mtp_start_loc, - draft_step, - accept_len, - lambda token_ids: token_ids, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py deleted file mode 100644 index 8827ba51d0..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_no_att.py +++ /dev/null @@ -1,72 +0,0 @@ -import copy - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal - - -class DpVanillaNoAttProposer(BaseDpProposer): - """普通 DP ``vanilla_no_att`` proposer。""" - - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - pass - - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> VanillaSpecProposal: - req_num = int(b_req_mtp_start_loc.shape[0]) - proposal_token_ids_by_step = [] - schedule_scores_by_step = [] - - if draft_step == 0: - return VanillaSpecProposal( - token_ids=target_next_token_ids.new_empty((req_num, 0)), - extra_mem_indexes_cpu=[], - schedule_scores=( - torch.empty((req_num, 0), dtype=torch.float32, device=target_next_token_ids.device) - if self.enable_dynmaic_mtp - else None - ), - ) - - assert accept_len is not None - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_token_ids = target_next_token_ids - draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_input = copy.copy(target_model_input) - - for step in range(draft_step): - draft_input.input_ids = draft_token_ids - draft_input.mtp_draft_input_hiddens = draft_hidden - draft_output = self.backend.draft_models[step].forward(draft_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - - if self.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) - selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) - schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) - else: - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) - proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) - - proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) - schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None - return VanillaSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=schedule_scores, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py deleted file mode 100644 index 30820c0ee7..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_proposers/vanilla_with_att.py +++ /dev/null @@ -1,156 +0,0 @@ -import copy - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids -from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager - - -class DpVanillaWithAttProposer(BaseDpProposer): - """普通 DP ``vanilla_with_att`` proposer。""" - - def fill_draft_model_kv_state( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - ) -> None: - assert target_model_input.is_prefill - assert target_model_input.b_position_delta is None - assert target_next_token_ids.shape == target_model_input.b_req_idx.shape - - draft_input = copy.copy(target_model_input) - draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_token_ids = target_next_token_ids - for draft_model in self.backend.draft_models: - draft_input = self._prepare_mtp_prefill_inputs( - model_input=draft_input, - b_next_token_ids=draft_token_ids, - mtp_draft_input_hiddens=draft_hidden, - ) - draft_output = draft_model.forward(draft_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - - def propose_next( - self, - target_model_input: ModelInput, - target_model_output: ModelOutput, - target_next_token_ids: torch.Tensor, - b_req_mtp_start_loc: torch.Tensor, - draft_step: int, - accept_len: torch.Tensor | None = None, - ) -> VanillaSpecProposal: - req_num = int(b_req_mtp_start_loc.shape[0]) - proposal_token_ids_by_step = [] - schedule_scores_by_step = [] - - assert not target_model_input.is_prefill - assert accept_len is not None - assert accept_len.shape == (req_num,) - assert draft_step == self.backend.max_draft_step - assert len(self.backend.draft_models) == draft_step - - accepted_tail_rows = (b_req_mtp_start_loc + accept_len - 1).long() - draft_token_ids = target_next_token_ids - draft_hidden = target_model_output.mtp_collector.spec_hidden - draft_input = copy.copy(target_model_input) - - for step in range(draft_step): - draft_input.input_ids = draft_token_ids - draft_input.mtp_draft_input_hiddens = draft_hidden - draft_output = self.backend.draft_models[step].forward(draft_input) - draft_hidden = draft_output.mtp_collector.spec_hidden - - if self.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) - selected_token_probs = draft_token_probs.index_select(0, accepted_tail_rows) - schedule_scores_by_step.append(selected_token_probs.float().unsqueeze(1)) - else: - draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) - selected_token_ids = draft_token_ids.index_select(0, accepted_tail_rows) - proposal_token_ids_by_step.append(selected_token_ids.unsqueeze(1)) - - if step + 1 < draft_step: - draft_token_ids = overlay_chained_mtp_decode_input( - input_ids=draft_input.input_ids, - draft_token_ids=draft_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - accept_len=accept_len, - ) - - proposal_token_ids = torch.cat(proposal_token_ids_by_step, dim=1) - schedule_scores = torch.cat(schedule_scores_by_step, dim=1) if self.enable_dynmaic_mtp else None - return VanillaSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=schedule_scores, - ) - - def _prepare_mtp_prefill_inputs( - self, - model_input: ModelInput, - b_next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: torch.Tensor, - ) -> ModelInput: - """构造 Vanilla chained MTP 下一层 draft model 的 prefill 输入。 - - 每个请求当前参与计算的 query 长度为 - ``b_seq_len - b_ready_cache_len``。本方法在各请求自己的 query - 区间内将 token 左移一位,即丢弃区间首 token,并在区间尾部追加该 - 请求的 ``b_next_token_ids``。这样,下一层 MTP model 看到的 token - 与上一层 model 生成的 next token 连续对齐。Chunked prefill 时只 - 移动当前 chunk,已经写入 KV cache 的 prefix 不参与移动。 - - 同时,本方法会完成两项辅助输入设置: - - 1. 将 ``b_is_decode_req`` 设置为全 False 的缓存 GPU 常量张量,确保 - mixed-prefill 逻辑保留这里显式生成的 token,而不会按 decode - 请求重新收集 token。 - 2. 将上一层 model 输出的 hidden states 绑定到 - ``mtp_draft_input_hiddens``,供下一层 MTP model 使用。 - - 参数: - model_input: 当前层的 prefill 输入。方法会原地更新该对象的 - ``input_ids``、``b_is_decode_req`` 和 - ``mtp_draft_input_hiddens`` 字段。 - b_next_token_ids: 每个请求需要追加到当前 query 尾部的 token, - shape 为 ``[batch_size]``。 - mtp_draft_input_hiddens: 上一层 model 产生、供下一层 draft model - 使用的 hidden states。 - - 返回: - 更新后的 ``model_input``,与传入对象是同一个对象。 - - 示例: - 假设 ``b_seq_len=[4, 5]``、``b_ready_cache_len=[1, 2]``,则两个 - 请求当前 query 长度均为 3。若扁平输入和追加 token 为: - - ``input_ids = [10, 11, 12, 20, 21, 22]`` - ``b_next_token_ids = [13, 23]`` - - 更新后的扁平输入为: - - ``input_ids = [11, 12, 13, 21, 22, 23]`` - - 两个请求已缓存的 prefix 长度分别为 1 和 2,它们对应的 token - 不在 ``input_ids`` 中,因此不会被本方法移动或重写。 - """ - model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( - key="dp_vanilla_mtp_prefill_b_is_decode_req", - shape=model_input.b_req_idx.shape, - fill_value=False, - dtype=torch.bool, - ) - model_input.input_ids = gen_mtp_new_input_ids( - input_ids=model_input.input_ids, - b_next_token_ids=b_next_token_ids, - b_seq_len=model_input.b_seq_len, - b_ready_cache_len=model_input.b_ready_cache_len, - ) - model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens - return model_input diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index a0088f9570..5b790ec726 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -21,14 +21,23 @@ class SpecEngine: - """Owns non-DP MTP planning and draft proposal generation. + """Owns MTP planning and draft proposal generation. Target verification, request metrics, stream synchronization, and resource - cleanup are stateless operations exposed by ``mtp_speculative.utils``. + cleanup are stateless operations exposed by ``mtp_speculative.utils``. DP + backends can enable ``allow_empty_batch`` so an idle rank still executes the + same fixed-depth draft-forward sequence as ranks that own real requests. """ - def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: + def __init__( + self, + backend: ModeBackend, + spec_mode: str, + enable_dynmaic_mtp: bool, + allow_empty_batch: bool = False, + ) -> None: self.backend = backend + self.allow_empty_batch = bool(allow_empty_batch) self.proposer: BaseSpecProposer = build_spec_proposer( spec_mode=spec_mode, backend=backend, @@ -58,7 +67,14 @@ def fill_draft_model_kv_state( def plan_decode(self, model_input: ModelInput, decode_reqs: List) -> SpecDecodePlan: """Return the fixed or dynamic speculative plan for one decode iteration.""" - assert decode_reqs, "non-DP speculative decode requires at least one request" + if not decode_reqs: + assert getattr(self, "allow_empty_batch", False), "speculative decode requires at least one request" + return SpecDecodePlan( + origin_batch_size=model_input.batch_size, + dynamic_batch_size=model_input.batch_size, + draft_step=self.backend.max_draft_step, + pre_draft_step=self.backend.max_draft_step, + ) return self.planner.plan( decode_reqs=decode_reqs, origin_batch_size=model_input.batch_size, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py index 27f790c54f..b9a379818c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -51,7 +51,8 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: Args: decode_reqs: 当前参与 decode 的非空逻辑请求列表。规划器可以读取 请求的输出进度,判断请求是否已经持有上一轮生成的 - draft proposal。空 batch 只由 DP 专用 planner 处理。 + draft proposal。DP 空 batch 由 SpecEngine 在进入 planner 前 + 直接构造固定深度计划。 origin_batch_size: 进入动态压缩前的物理 verify 行数。 Returns: diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index ac7efc2b45..5bd9908db0 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -1,14 +1,15 @@ from types import SimpleNamespace +import pytest import torch from lightllm.server.router.model_infer.mode_backend.dp_backend import impl as dp_backend_impl from lightllm.server.router.model_infer.mode_backend.dp_backend.impl import DPChunkedPrefillBackend from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ChunkedPrefillBackend -from lightllm.server.router.model_infer.mtp_speculative.dp_engine import DPSpecEngine from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine +from lightllm.server.router.model_infer.mtp_speculative.planner import FixedSpecPlanner, SpecDecodePlan from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( MtpMemIndexesToFree, SpecProposal, @@ -52,7 +53,7 @@ def _assert_all_mem_indexes_are_freed(extra_mem_indexes_cpu, expected): assert extra_mem_indexes_cpu[0].free_mask_cpu is None -def test_backends_initialize_their_own_spec_engine(): +def test_dp_backend_reuses_common_engine_outside_overlap(): args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False) backend = ChunkedPrefillBackend.__new__(ChunkedPrefillBackend) backend.args = args @@ -67,11 +68,13 @@ def test_backends_initialize_their_own_spec_engine(): assert "spec_engine_class" not in ModeBackend.__dict__ assert type(backend.spec_engine) is SpecEngine - assert type(dp_backend.spec_engine) is DPSpecEngine + assert type(dp_backend.spec_engine) is SpecEngine + assert dp_backend.spec_engine.allow_empty_batch + assert not dp_backend.spec_engine.proposer.enable_dynmaic_mtp + assert type(dp_backend.spec_engine.planner) is FixedSpecPlanner assert type(dp_backend.dp_overlap_spec_engine) is DPOverlapSpecEngine assert dp_backend.prefill_draft_engine is dp_backend.spec_engine assert dp_backend.decode_draft_engine is dp_backend.dp_overlap_spec_engine - assert not issubclass(DPSpecEngine, SpecEngine) assert not issubclass(DPOverlapSpecEngine, SpecEngine) @@ -104,67 +107,62 @@ def test_padded_token_ids_support_empty_dp_rank(): assert torch.equal(padded_token_ids, torch.zeros(4, dtype=torch.int64)) -def test_dp_eagle_uses_common_extend_then_unit_decode_proposer(monkeypatch): - scatter_args = _capture_scatter_args(monkeypatch) - backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) - backend.max_draft_step = 7 - backend.spec_engine = _RecordingSpecEngine() - backend.decode_draft_engine = backend.spec_engine +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_dp_decode_mtp_runs_common_engine_for_empty_batch(monkeypatch): + device = "cuda" + empty_i32 = torch.empty((0,), dtype=torch.int32, device=device) model_input = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), - ) - model_output = SimpleNamespace(spec_hidden=torch.randn(16, 4)) - next_token_ids = torch.arange(8, dtype=torch.int64) - real_start_locs = torch.tensor([0], dtype=torch.int32) - real_accept_len = torch.tensor([2], dtype=torch.int32) - - extra_mem = backend._draft_decode_eagle( - model_input=model_input, - model_output=model_output, - next_token_ids=next_token_ids, - b_req_mtp_start_loc=real_start_locs, - mtp_accept_len=real_accept_len, - req_num=8, + batch_size=0, + b_req_idx=empty_i32, + b_mtp_index=empty_i32, + mem_indexes_cpu=torch.empty((0,), dtype=torch.int32), ) + model_output = SimpleNamespace(logits=torch.empty((0, 8), device=device)) + calls = [] - propose_args = backend.spec_engine.propose_args - assert propose_args["target_next_token_ids"].shape == (16,) - assert torch.equal(propose_args["target_next_token_ids"][:8], next_token_ids) - assert torch.equal(propose_args["b_req_mtp_start_loc"], torch.tensor([0, 8], dtype=torch.int32)) - assert torch.equal(propose_args["accept_len"], torch.tensor([2, 1], dtype=torch.int32)) - assert scatter_args["proposal"].token_ids.shape == (1, 7) - assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) - _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([123], dtype=torch.int32)) + class _CommonEngine: + def plan_decode(self, **kwargs): + calls.append("plan") + return SpecDecodePlan(0, 0, 2, 2) + def prepare_decode_model_input(self, **kwargs): + calls.append("prepare") + return kwargs["model_input"], None + + def propose_next(self, **kwargs): + calls.append("propose") + assert kwargs["target_next_token_ids"].shape == (0,) + assert kwargs["b_req_mtp_start_loc"].shape == (0,) + assert kwargs["accept_len"].shape == (0,) + return SpecProposal(token_ids=torch.empty((0, 2), dtype=torch.int64, device=device)) -def test_dp_vanilla_uses_dp_engine_proposer(monkeypatch): - scatter_args = _capture_scatter_args(monkeypatch) backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) - backend.max_draft_step = 7 - backend.spec_engine = _RecordingSpecEngine() - backend.decode_draft_engine = backend.spec_engine - model_input = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), + backend.spec_engine = _CommonEngine() + backend.model = SimpleNamespace(forward=lambda _: model_output) + event_pack = SimpleNamespace( + notify_post_handle_and_wait_pre_post_handle=lambda: calls.append("post_wait"), + notify_forward_and_wait_post_handle=lambda: calls.append("forward_wait"), + notify_pre_post_handle=lambda: calls.append("pre_post"), ) - next_token_ids = torch.arange(8, dtype=torch.int64) - - extra_mem = backend._draft_decode_vanilla( - model_input=model_input, - model_output=SimpleNamespace(), - next_token_ids=next_token_ids, - b_req_mtp_start_loc=torch.arange(8, dtype=torch.int32), - mtp_accept_len=torch.ones(8, dtype=torch.int32), - req_num=8, + monkeypatch.setattr( + dp_backend_impl, + "prepare_decode_inputs", + lambda req_objs: (model_input, []), + ) + monkeypatch.setattr( + dp_backend_impl, + "g_infer_context", + SimpleNamespace(get_overlap_stream=lambda: torch.cuda.current_stream()), + ) + monkeypatch.setattr( + dp_backend_impl.mtp_utils, + "free_mem_indexes", + lambda **kwargs: calls.append("free"), ) - propose_args = backend.spec_engine.propose_args - assert propose_args["target_next_token_ids"].shape == (16,) - assert torch.equal(propose_args["target_next_token_ids"][:8], next_token_ids) - assert scatter_args["proposal"].token_ids.shape == (8, 7) - assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) - _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([123], dtype=torch.int32)) + backend.decode_mtp(event_pack=event_pack, decode_reqs=[]) + + assert calls == ["plan", "prepare", "propose", "post_wait", "forward_wait", "free", "pre_post"] def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(monkeypatch): diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 737412d408..ebebe27612 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -5,25 +5,11 @@ import torch from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine -from lightllm.server.router.model_infer.mtp_speculative.dp_planner import ( - BaseDpPlanner, - FixedDpPlanner, - build_dp_planner, -) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner import ( BaseDpOverlapPlanner, FixedDpOverlapPlanner, build_dp_overlap_planner, ) -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers import build_dp_spec_proposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.base import BaseDpProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle3 import DpEagle3Proposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_no_att import DpEagleNoAttProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.eagle_with_att import DpEagleWithAttProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_no_att import DpVanillaNoAttProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_with_att import ( - DpVanillaWithAttProposer, -) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers import ( build_dp_overlap_spec_proposer, ) @@ -135,6 +121,21 @@ def test_non_dp_engine_rejects_empty_decode_batch(): engine.plan_decode(model_input=SimpleNamespace(batch_size=0), decode_reqs=[]) +def test_common_engine_builds_fixed_plan_for_allowed_empty_batch(): + engine = SpecEngine.__new__(SpecEngine) + engine.allow_empty_batch = True + engine.backend = SimpleNamespace(max_draft_step=3) + + plan = engine.plan_decode(model_input=SimpleNamespace(batch_size=0), decode_reqs=[]) + + assert plan == SpecDecodePlan( + origin_batch_size=0, + dynamic_batch_size=0, + draft_step=3, + pre_draft_step=3, + ) + + def test_spec_engine_only_exposes_planning_and_proposal_interfaces(): public_methods = { name for name, value in SpecEngine.__dict__.items() if callable(value) and not name.startswith("_") @@ -244,14 +245,6 @@ def test_scatter_mtp_next_tokens_ignores_empty_schedule_scores(monkeypatch): assert scatter_args["schedule_scores"] is None -def test_dp_planner_returns_fixed_backend_draft_step(): - planner = build_dp_planner(backend=SimpleNamespace(max_draft_step=4)) - - assert type(planner) is FixedDpPlanner - assert isinstance(planner, BaseDpPlanner) - assert planner.get_draft_step() == 4 - - def test_dp_overlap_planner_returns_fixed_backend_draft_step(): planner = build_dp_overlap_planner(backend=SimpleNamespace(max_draft_step=5)) @@ -323,13 +316,6 @@ def test_each_mode_proposer_inherits_its_expected_implementation_base(): DFlashProposer, DSparkProposer, ) - dp_proposer_types = ( - DpVanillaWithAttProposer, - DpVanillaNoAttProposer, - DpEagleWithAttProposer, - DpEagleNoAttProposer, - DpEagle3Proposer, - ) dp_overlap_proposer_types = ( DpOverlapVanillaWithAttProposer, DpOverlapVanillaNoAttProposer, @@ -341,8 +327,6 @@ def test_each_mode_proposer_inherits_its_expected_implementation_base(): for proposer_type in proposer_types: assert proposer_type.__bases__ == (BaseSpecProposer,) assert Eagle3Proposer.__bases__ == (EagleWithAttProposer,) - for proposer_type in dp_proposer_types: - assert proposer_type.__bases__ == (BaseDpProposer,) for proposer_type in dp_overlap_proposer_types: assert proposer_type.__bases__ == (BaseDpOverlapProposer,) @@ -364,22 +348,6 @@ def test_each_mtp_mode_builds_its_own_proposer(): assert type(proposer) is proposer_type -def test_each_supported_dp_mtp_mode_builds_its_own_dp_proposer(): - backend = SimpleNamespace() - proposer_types = { - "vanilla_with_att": DpVanillaWithAttProposer, - "vanilla_no_att": DpVanillaNoAttProposer, - "eagle_with_att": DpEagleWithAttProposer, - "eagle_no_att": DpEagleNoAttProposer, - "eagle3": DpEagle3Proposer, - } - - for spec_mode, proposer_type in proposer_types.items(): - proposer = build_dp_spec_proposer(spec_mode=spec_mode, backend=backend, enable_dynmaic_mtp=False) - assert type(proposer) is proposer_type - assert isinstance(proposer, BaseDpProposer) - - def test_each_supported_dp_overlap_mtp_mode_builds_its_own_proposer(): backend = SimpleNamespace() proposer_types = { diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py index 7f0f7a442f..8ce9114aab 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py @@ -7,7 +7,6 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_no_att import ( DpOverlapVanillaNoAttProposer, ) -from lightllm.server.router.model_infer.mtp_speculative.dp_proposers.vanilla_no_att import DpVanillaNoAttProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer @@ -58,6 +57,57 @@ def test_select_accepted_tail_rows_triton_matches_index_select(): torch.testing.assert_close(selected.b_position_delta, b_position_delta.index_select(0, expected_rows)) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_vanilla_no_att_runs_all_draft_steps_for_empty_dp_batch(): + device = "cuda" + draft_batch_sizes = [] + + def draft_forward(model_input): + draft_batch_sizes.append(model_input.batch_size) + return SimpleNamespace( + token_ids=torch.empty((0,), dtype=torch.int64, device=device), + mtp_collector=SimpleNamespace( + spec_hidden=torch.empty((0, 2), dtype=torch.float32, device=device), + ), + ) + + backend = SimpleNamespace( + draft_models=[SimpleNamespace(forward=draft_forward) for _ in range(2)], + _gen_argmax_token_ids=lambda output: output.token_ids, + ) + proposer = VanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False) + empty_i32 = torch.empty((0,), dtype=torch.int32, device=device) + target_model_input = SimpleNamespace( + batch_size=0, + b_req_idx=empty_i32, + b_mtp_index=empty_i32, + b_seq_len=empty_i32, + mem_indexes=empty_i32, + mem_indexes_cpu=torch.empty((0,), dtype=torch.int32), + b_position_delta=empty_i32, + b_shared_seq_len=empty_i32, + b_shared_radix_node_id=torch.empty((0,), dtype=torch.int64, device=device), + multimodal_params=[], + ) + target_model_output = SimpleNamespace( + mtp_collector=SimpleNamespace( + spec_hidden=torch.empty((0, 2), dtype=torch.float32, device=device), + ) + ) + + proposal = proposer.propose_next( + target_model_input=target_model_input, + target_model_output=target_model_output, + target_next_token_ids=torch.empty((0,), dtype=torch.int64, device=device), + b_req_mtp_start_loc=empty_i32, + draft_step=2, + accept_len=empty_i32, + ) + + assert draft_batch_sizes == [0, 0] + assert proposal.token_ids.shape == (0, 2) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_vanilla_no_att_proposes_from_one_accepted_tail_per_request(): device = "cuda" @@ -161,15 +211,10 @@ def test_vanilla_no_att_skips_draft_forward_for_zero_steps(): assert proposal.schedule_scores.shape == (2, 0) -def test_all_vanilla_no_att_fill_hooks_are_noops(): +def test_vanilla_no_att_fill_hooks_are_noops(): backend = SimpleNamespace(draft_models=[]) - proposers = [ - VanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False), - DpVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False), - DpOverlapVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False), - ] - - for proposer in proposers: - proposer.fill_draft_model_kv_state(None, None, None) + proposer = VanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False) + overlap_proposer = DpOverlapVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False) - proposers[-1].fill_draft_model_kv_state_overlap(None, None, None, None, None, None) + proposer.fill_draft_model_kv_state(None, None, None) + overlap_proposer.fill_draft_model_kv_state_overlap(None, None, None, None, None, None) From 4398e9adea203dcb22b7656a0affd599545c9714 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 21 Aug 2026 12:36:08 +0000 Subject: [PATCH 085/103] fix: synchronize dynamic LightSpec draft steps across dp --- .../mode_backend/dp_backend/impl.py | 8 +- .../model_infer/mtp_speculative/engine.py | 14 +- .../mtp_speculative/planner/base.py | 5 +- .../mtp_speculative/planner/lightspec.py | 75 +++++++-- .../test_dp_overlap_spec_engine.py | 158 +++++++++++++++++- .../mtp_speculative/test_planner.py | 34 +++- 6 files changed, 250 insertions(+), 44 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 6d34dba034..2ee2f550c4 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -73,15 +73,13 @@ def init_spec_engine(self): spec_mode=self.args.mtp_mode, enable_dynmaic_mtp=self.args.mtp_dynamic_verify, ) - # 非 overlap DP 与普通后端复用同一个 SpecEngine。固定 draft 深度 - # 保证所有 DP rank 执行相同数量的 draft forward;空 rank 则由 - # allow_empty_batch 进入完整的 dummy forward 流程。 + # 非 overlap DP 与普通后端复用同一个 SpecEngine。 self.spec_engine = SpecEngine( backend=self, spec_mode=self.args.mtp_mode, - enable_dynmaic_mtp=False, - allow_empty_batch=True, + enable_dynmaic_mtp=self.args.mtp_dynamic_verify, ) + self.dp_overlap_spec_engine = DPOverlapSpecEngine(**engine_kwargs) self.prefill_draft_engine = ( self.dp_overlap_spec_engine if self.enable_prefill_microbatch_overlap else self.spec_engine diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 5b790ec726..9d59afd8a7 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -24,9 +24,7 @@ class SpecEngine: """Owns MTP planning and draft proposal generation. Target verification, request metrics, stream synchronization, and resource - cleanup are stateless operations exposed by ``mtp_speculative.utils``. DP - backends can enable ``allow_empty_batch`` so an idle rank still executes the - same fixed-depth draft-forward sequence as ranks that own real requests. + cleanup are stateless operations exposed by ``mtp_speculative.utils``. """ def __init__( @@ -34,10 +32,8 @@ def __init__( backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: bool, - allow_empty_batch: bool = False, ) -> None: self.backend = backend - self.allow_empty_batch = bool(allow_empty_batch) self.proposer: BaseSpecProposer = build_spec_proposer( spec_mode=spec_mode, backend=backend, @@ -67,14 +63,6 @@ def fill_draft_model_kv_state( def plan_decode(self, model_input: ModelInput, decode_reqs: List) -> SpecDecodePlan: """Return the fixed or dynamic speculative plan for one decode iteration.""" - if not decode_reqs: - assert getattr(self, "allow_empty_batch", False), "speculative decode requires at least one request" - return SpecDecodePlan( - origin_batch_size=model_input.batch_size, - dynamic_batch_size=model_input.batch_size, - draft_step=self.backend.max_draft_step, - pre_draft_step=self.backend.max_draft_step, - ) return self.planner.plan( decode_reqs=decode_reqs, origin_batch_size=model_input.batch_size, diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py index b9a379818c..c1becbdcd2 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -49,10 +49,9 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: """为当前 decode 迭代生成执行计划。 Args: - decode_reqs: 当前参与 decode 的非空逻辑请求列表。规划器可以读取 + decode_reqs: 当前参与 decode 的逻辑请求列表。规划器可以读取 请求的输出进度,判断请求是否已经持有上一轮生成的 - draft proposal。DP 空 batch 由 SpecEngine 在进入 planner 前 - 直接构造固定深度计划。 + draft proposal;DP 空 rank 对应空列表。 origin_batch_size: 进入动态压缩前的物理 verify 行数。 Returns: diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index 24cfbc78b8..2cd781ce49 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -3,6 +3,8 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Tuple import numpy as np +import torch +import torch.distributed as dist from lightllm.server.router.model_infer.mtp_speculative.planner.base import ( BaseMtpPlanner, @@ -10,6 +12,7 @@ _EMAValue, _InferCostMsTable, ) +from lightllm.utils.dist_utils import get_global_world_size if TYPE_CHECKING: from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend @@ -65,10 +68,32 @@ def __init__( # The current verify width is bounded by the proposal built last time. self.pre_draft_step = self.max_draft_step + # 只有非 overlap DP 下的变长 LightSpec 需要跨 rank 对齐 draft 深度。 + self._draft_step_group = None + self._draft_step_tensor = None + self._draft_step_stream = None + if backend.args.dp > 1 and not backend.enable_decode_microbatch_overlap and len(self.draft_steps) > 1: + self._draft_step_group = dist.new_group( + ranks=list(range(get_global_world_size())), + backend="nccl", + ) + self._draft_step_tensor = torch.zeros((1,), dtype=torch.int32, device="cuda") + self._draft_step_stream = torch.cuda.Stream() + def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: req_num = len(decode_reqs) pre_draft_step = self.pre_draft_step + if req_num == 0: + return self._sync_draft_step( + SpecDecodePlan( + origin_batch_size=origin_batch_size, + dynamic_batch_size=origin_batch_size, + draft_step=self.draft_steps[0], + pre_draft_step=pre_draft_step, + ) + ) + # A request entering its first decode has only the guaranteed target # row. Existing requests additionally expose pre_draft_step candidates. # This bounds Verify by the candidate pool that physically exists. @@ -81,13 +106,14 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: # Costs come from CUDA Graph capture, while progress requires one # real verification batch. Keep the available proposal intact until # J(N, B, d) has a progress observation. - self.pre_draft_step = self.max_draft_step - return SpecDecodePlan( - origin_batch_size=origin_batch_size, - dynamic_batch_size=max_batch_size, - draft_step=self.max_draft_step, - pre_draft_step=pre_draft_step, - all_reqs_have_proposals=all_reqs_have_proposals, + return self._sync_draft_step( + SpecDecodePlan( + origin_batch_size=origin_batch_size, + dynamic_batch_size=max_batch_size, + draft_step=self.max_draft_step, + pre_draft_step=pre_draft_step, + all_reqs_have_proposals=all_reqs_have_proposals, + ) ) min_batch_size = req_num @@ -103,13 +129,40 @@ def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: dynamic_batch_size = batch_sizes[np.argmin(costs)] draft_step = self._select_draft_step(req_num=req_num) + return self._sync_draft_step( + SpecDecodePlan( + origin_batch_size=origin_batch_size, + dynamic_batch_size=dynamic_batch_size, + draft_step=draft_step, + pre_draft_step=pre_draft_step, + all_reqs_have_proposals=all_reqs_have_proposals, + ) + ) + + def _sync_draft_step(self, plan: SpecDecodePlan) -> SpecDecodePlan: + """在 LightSpec 的全局 NCCL 组内对齐下一轮 draft 深度。""" + + draft_step = int(plan.draft_step) + if self._draft_step_group is not None: + assert self._draft_step_stream is not None + with torch.cuda.stream(self._draft_step_stream): + self._draft_step_tensor.fill_(draft_step) + dist.all_reduce( + self._draft_step_tensor, + op=dist.ReduceOp.MAX, + group=self._draft_step_group, + async_op=False, + ) + draft_step = int(self._draft_step_tensor.item()) + + assert draft_step in self.draft_steps self.pre_draft_step = draft_step return SpecDecodePlan( - origin_batch_size=origin_batch_size, - dynamic_batch_size=dynamic_batch_size, + origin_batch_size=plan.origin_batch_size, + dynamic_batch_size=plan.dynamic_batch_size, draft_step=draft_step, - pre_draft_step=pre_draft_step, - all_reqs_have_proposals=all_reqs_have_proposals, + pre_draft_step=plan.pre_draft_step, + all_reqs_have_proposals=plan.all_reqs_have_proposals, ) def update_statics( diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 5bd9908db0..1ede7c4acb 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -1,3 +1,4 @@ +from contextlib import nullcontext from types import SimpleNamespace import pytest @@ -9,7 +10,12 @@ from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ChunkedPrefillBackend from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine -from lightllm.server.router.model_infer.mtp_speculative.planner import FixedSpecPlanner, SpecDecodePlan +from lightllm.server.router.model_infer.mtp_speculative.planner import lightspec as lightspec_planner_impl +from lightllm.server.router.model_infer.mtp_speculative.planner import ( + FixedSpecPlanner, + LightSpecPlanner, + SpecDecodePlan, +) from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( MtpMemIndexesToFree, SpecProposal, @@ -54,7 +60,7 @@ def _assert_all_mem_indexes_are_freed(extra_mem_indexes_cpu, expected): def test_dp_backend_reuses_common_engine_outside_overlap(): - args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False) + args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False, dp=1) backend = ChunkedPrefillBackend.__new__(ChunkedPrefillBackend) backend.args = args backend.max_draft_step = 2 @@ -62,6 +68,7 @@ def test_dp_backend_reuses_common_engine_outside_overlap(): dp_backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) dp_backend.args = args dp_backend.max_draft_step = 2 + dp_backend.dp_size = 1 dp_backend.enable_prefill_microbatch_overlap = False dp_backend.enable_decode_microbatch_overlap = True dp_backend.init_spec_engine() @@ -69,7 +76,6 @@ def test_dp_backend_reuses_common_engine_outside_overlap(): assert "spec_engine_class" not in ModeBackend.__dict__ assert type(backend.spec_engine) is SpecEngine assert type(dp_backend.spec_engine) is SpecEngine - assert dp_backend.spec_engine.allow_empty_batch assert not dp_backend.spec_engine.proposer.enable_dynmaic_mtp assert type(dp_backend.spec_engine.planner) is FixedSpecPlanner assert type(dp_backend.dp_overlap_spec_engine) is DPOverlapSpecEngine @@ -78,14 +84,158 @@ def test_dp_backend_reuses_common_engine_outside_overlap(): assert not issubclass(DPOverlapSpecEngine, SpecEngine) +def test_lightspec_planner_reduces_draft_step_with_max(monkeypatch): + group = object() + planner = LightSpecPlanner.__new__(LightSpecPlanner) + planner._draft_step_group = group + planner._draft_step_tensor = torch.zeros((1,), dtype=torch.int32) + planner._draft_step_stream = object() + planner.draft_steps = (1, 2, 3) + planner.pre_draft_step = 1 + planner.backend = SimpleNamespace(dp_size=2) + + def all_reduce(tensor, op, group, async_op): + assert op == torch.distributed.ReduceOp.MAX + assert group is planner._draft_step_group + assert not async_op + tensor.fill_(3) + + monkeypatch.setattr(lightspec_planner_impl.dist, "all_reduce", all_reduce) + monkeypatch.setattr(lightspec_planner_impl.torch.cuda, "stream", lambda stream: nullcontext()) + + plan = planner.plan(decode_reqs=[], origin_batch_size=2) + + assert plan.dynamic_batch_size == 2 + assert plan.draft_step == 3 + assert planner.pre_draft_step == 3 + + +def test_dp_backend_builds_dedicated_global_nccl_group(monkeypatch): + created_group = object() + created_stream = object() + new_group_args = {} + real_torch_zeros = torch.zeros + + def new_group(*, ranks, backend): + new_group_args.update(ranks=ranks, backend=backend) + return created_group + + def zeros_without_cuda(*args, **kwargs): + kwargs.pop("device", None) + return real_torch_zeros(*args, **kwargs) + + monkeypatch.setattr(lightspec_planner_impl, "get_global_world_size", lambda: 4) + monkeypatch.setattr(lightspec_planner_impl.dist, "new_group", new_group) + monkeypatch.setattr(lightspec_planner_impl.torch, "zeros", zeros_without_cuda) + monkeypatch.setattr(lightspec_planner_impl.torch.cuda, "Stream", lambda: created_stream) + + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=True, dp=2) + backend.max_draft_step = 3 + backend.dp_size = 2 + backend.model = SimpleNamespace(graph=None) + backend.draft_models = [SimpleNamespace(graph=None)] + backend.enable_prefill_microbatch_overlap = False + backend.enable_decode_microbatch_overlap = False + backend.init_spec_engine() + + assert new_group_args == {"ranks": [0, 1, 2, 3], "backend": "nccl"} + planner = backend.spec_engine.planner + assert planner._draft_step_group is created_group + assert planner._draft_step_tensor.shape == (1,) + assert planner._draft_step_stream is created_stream + + +def test_dp_backend_does_not_build_group_for_single_draft_step_mode(monkeypatch): + monkeypatch.setattr( + lightspec_planner_impl.dist, + "new_group", + lambda **kwargs: pytest.fail(f"unexpected NCCL group: {kwargs}"), + ) + + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.args = SimpleNamespace(mtp_mode="vanilla_with_att", mtp_dynamic_verify=True, dp=2) + backend.max_draft_step = 3 + backend.dp_size = 2 + backend.model = SimpleNamespace(graph=None) + backend.draft_models = [SimpleNamespace(graph=None, block_size=4)] + backend.enable_prefill_microbatch_overlap = False + backend.enable_decode_microbatch_overlap = False + + backend.init_spec_engine() + + assert backend.spec_engine.planner._draft_step_group is None + assert backend.spec_engine.planner._draft_step_tensor is None + assert backend.spec_engine.planner._draft_step_stream is None + + +def test_dp_backend_does_not_build_group_for_fixed_planner(monkeypatch): + monkeypatch.setattr( + lightspec_planner_impl.dist, + "new_group", + lambda **kwargs: pytest.fail(f"unexpected NCCL group: {kwargs}"), + ) + + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False, dp=2) + backend.max_draft_step = 3 + backend.dp_size = 2 + backend.enable_prefill_microbatch_overlap = False + backend.enable_decode_microbatch_overlap = False + + backend.init_spec_engine() + + assert not hasattr(backend.spec_engine.planner, "_draft_step_group") + + +def test_dp_backend_does_not_build_group_for_overlap_decode(monkeypatch): + monkeypatch.setattr( + lightspec_planner_impl.dist, + "new_group", + lambda **kwargs: pytest.fail(f"unexpected NCCL group: {kwargs}"), + ) + + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=True, dp=2) + backend.max_draft_step = 3 + backend.dp_size = 2 + backend.model = SimpleNamespace(graph=None) + backend.draft_models = [SimpleNamespace(graph=None)] + backend.enable_prefill_microbatch_overlap = False + backend.enable_decode_microbatch_overlap = True + + backend.init_spec_engine() + + assert backend.spec_engine.planner._draft_step_group is None + assert backend.decode_draft_engine is backend.dp_overlap_spec_engine + + +def test_dp_backend_keeps_dynamic_planning_before_global_reduction(): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=True, dp=1) + backend.max_draft_step = 3 + backend.dp_size = 1 + backend.model = SimpleNamespace(graph=None) + backend.draft_models = [SimpleNamespace(graph=None)] + backend.enable_prefill_microbatch_overlap = False + backend.enable_decode_microbatch_overlap = False + + backend.init_spec_engine() + + assert isinstance(backend.spec_engine.planner, LightSpecPlanner) + assert backend.spec_engine.proposer.enable_dynmaic_mtp + assert backend.spec_engine.planner._draft_step_group is None + + def test_dp_prefill_and_decode_select_overlap_engine_independently(): - args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False) + args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False, dp=1) for prefill_overlap in (False, True): for decode_overlap in (False, True): backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) backend.args = args backend.max_draft_step = 2 + backend.dp_size = 1 backend.enable_prefill_microbatch_overlap = prefill_overlap backend.enable_decode_microbatch_overlap = decode_overlap backend.init_spec_engine() diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index ebebe27612..3f511ab0ac 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -63,6 +63,8 @@ def build_lightspec_planner( block_size: int = 3, ): backend = SimpleNamespace( + args=SimpleNamespace(dp=1), + enable_decode_microbatch_overlap=False, max_draft_step=max_draft_step, model=SimpleNamespace(graph=None), draft_models=[SimpleNamespace(block_size=block_size, graph=None)], @@ -93,6 +95,8 @@ def build_decode_reqs(req_num: int, req_num_with_proposals: int | None = None): def build_planner(spec_mode: str, enable_dynmaic_mtp: bool = True): engine = SpecEngine.__new__(SpecEngine) engine.backend = SimpleNamespace( + args=SimpleNamespace(dp=1), + enable_decode_microbatch_overlap=False, max_draft_step=3, model=SimpleNamespace(graph=None), draft_models=[SimpleNamespace(block_size=3, graph=None)], @@ -114,25 +118,34 @@ def test_fixed_planner_returns_static_plan(): assert not plan.skip_verify_sync -def test_non_dp_engine_rejects_empty_decode_batch(): +def test_common_engine_accepts_empty_fixed_decode_batch(): engine = SpecEngine.__new__(SpecEngine) + engine.planner = FixedSpecPlanner(max_draft_step=3) - with pytest.raises(AssertionError, match="requires at least one request"): - engine.plan_decode(model_input=SimpleNamespace(batch_size=0), decode_reqs=[]) + plan = engine.plan_decode(model_input=SimpleNamespace(batch_size=0), decode_reqs=[]) + + assert plan == SpecDecodePlan( + origin_batch_size=0, + dynamic_batch_size=0, + draft_step=3, + pre_draft_step=3, + ) -def test_common_engine_builds_fixed_plan_for_allowed_empty_batch(): +def test_common_engine_delegates_empty_dp_batch_to_lightspec_planner(): engine = SpecEngine.__new__(SpecEngine) - engine.allow_empty_batch = True - engine.backend = SimpleNamespace(max_draft_step=3) + engine.backend = SimpleNamespace(max_draft_step=3, dp_size=2) + engine.planner = build_lightspec_planner(spec_mode="eagle3") + engine.planner.backend.args.dp = 2 + engine.planner.pre_draft_step = 2 plan = engine.plan_decode(model_input=SimpleNamespace(batch_size=0), decode_reqs=[]) assert plan == SpecDecodePlan( origin_batch_size=0, dynamic_batch_size=0, - draft_step=3, - pre_draft_step=3, + draft_step=1, + pre_draft_step=2, ) @@ -288,9 +301,14 @@ def test_engine_routes_only_dspark_to_the_confidence_planner(): assert isinstance(eagle_planner, LightSpecPlanner) assert eagle_planner.draft_steps == (1, 2, 3) + eagle_with_att_planner = build_planner("eagle_with_att") + assert eagle_with_att_planner.draft_steps == (1, 2, 3) + def test_dynamic_planner_registers_cuda_graph_costs_from_backend(): backend = SimpleNamespace( + args=SimpleNamespace(dp=1), + enable_decode_microbatch_overlap=False, max_draft_step=3, model=SimpleNamespace(graph=SimpleNamespace(infer_cost_ms_by_batch_size={4: 1.2})), draft_models=[ From b9957ab0abcda0ea0df24c2aa12dc3500d940963 Mon Sep 17 00:00:00 2001 From: sufubao Date: Fri, 21 Aug 2026 20:54:07 +0800 Subject: [PATCH 086/103] chore: add DSpark 1P1D deployment script --- deploy_dspark_1p1d.sh | 175 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 175 insertions(+) create mode 100755 deploy_dspark_1p1d.sh diff --git a/deploy_dspark_1p1d.sh b/deploy_dspark_1p1d.sh new file mode 100755 index 0000000000..4eadf20f19 --- /dev/null +++ b/deploy_dspark_1p1d.sh @@ -0,0 +1,175 @@ +#!/usr/bin/env bash +set -euo pipefail + +PORT="${PORT:-16666}" +MODEL_DIR="${MODEL_DIR:-/mnt/afs/models/Qwen3.5-27B}" +DRAFT_MODEL_DIR="${DRAFT_MODEL_DIR:-${draft_model:-}}" +MODEL_NAME="${MODEL_NAME:-sensenova-flash-lite-v41-20260816-fp8-step3k-dspark}" + +CPU_CACHE_SIZE="${CPU_CACHE_SIZE:-600}" +CHAT_TEMPLATE="${CHAT_TEMPLATE:-${MODEL_DIR}/chat_template.jinja}" +PD_MASTER_IP="${PD_MASTER_IP:-127.0.0.1}" + +# DSpark checkpoint block_size. Prefill and decode must use the same value. +MTP_STEP="${MTP_STEP:-5}" + +# DSpark expands each logical decode request into speculative verification rows. +# Start conservatively; raise both values together after checking GPU headroom. +D_MAX_REQ_SIZE="${D_MAX_REQ_SIZE:-16}" +D_GRAPH_MAX_BATCH_SIZE="${D_GRAPH_MAX_BATCH_SIZE:-16}" + +if [[ -z "${DRAFT_MODEL_DIR}" ]]; then + echo "DRAFT_MODEL_DIR is required (the DSpark draft checkpoint directory)." >&2 + exit 2 +fi + +export LOADWORKER="${LOADWORKER:-8}" +export LIGHTLLM_TRITON_AUTOTUNE_LEVEL="${LIGHTLLM_TRITON_AUTOTUNE_LEVEL:-1}" +export LIGHTLLM_ANTHROPIC_ENABLE_PDF_PARSING="${LIGHTLLM_ANTHROPIC_ENABLE_PDF_PARSING:-1}" +export LIGHTLLM_LOG_LEVEL="${LIGHTLLM_LOG_LEVEL:-debug}" +export PYTHONUNBUFFERED=1 +export LIGHTLLM_PD_SPLIT_MAX_NEW_TOKENS="${LIGHTLLM_PD_SPLIT_MAX_NEW_TOKENS:-4096}" + +# Keep same-host PD WebSocket traffic away from HTTP/WebSocket proxies. +export NO_PROXY="${NO_PROXY:+${NO_PROXY},}${PD_MASTER_IP},localhost" +export no_proxy="${NO_PROXY}" + +RUN_NAME="${RUN_NAME:-qwen35-dspark-1p1d}" +LOG_DIR="${LOG_ROOT:-/mnt/afs/lightllm-runs}/${RUN_NAME}" +mkdir -p "${LOG_DIR}" + +P_COMMON_ARGS=( + --model_dir "${MODEL_DIR}" + --model_name "${MODEL_NAME}" + --graph_max_batch_size 32 + --running_max_req_size 32 + --mem_fraction 0.75 + --max_image_token_count 4096 + --max_image_pixels 3686400 + --batch_max_tokens 32768 + --chunked_prefill_size 16384 + --linear_att_cache_size 500 + --linear_att_hash_page_size 2048 + --linear_att_page_block_num 8 + --quant_type fp8w8a8-pt-sgl + --mtp_mode dspark + --mtp_draft_model_dir "${DRAFT_MODEL_DIR}" + --mtp_step "${MTP_STEP}" + --chat_template "${CHAT_TEMPLATE}" + --pd_trans_mode nixl + --pd_kv_page_size 4096 + --pd_master_ip "${PD_MASTER_IP}" + --pd_master_port "${PORT}" + --llm_prefill_att_backend fa3 flashqla +) + +D_COMMON_ARGS=( + --model_dir "${MODEL_DIR}" + --model_name "${MODEL_NAME}" + --graph_max_batch_size "${D_GRAPH_MAX_BATCH_SIZE}" + --running_max_req_size "${D_MAX_REQ_SIZE}" + --mem_fraction 0.75 + --max_image_token_count 4096 + --max_image_pixels 3686400 + --batch_max_tokens 256 + --linear_att_cache_size 500 + --quant_type fp8w8a8-pt-sgl + --mtp_mode dspark + --mtp_draft_model_dir "${DRAFT_MODEL_DIR}" + --mtp_step "${MTP_STEP}" + --chat_template "${CHAT_TEMPLATE}" + --pd_trans_mode nixl + --pd_kv_page_size 4096 + --pd_master_ip "${PD_MASTER_IP}" + --pd_master_port "${PORT}" +) + +PIDS=() + +cleanup() { + if ((${#PIDS[@]} > 0)); then + kill -TERM "${PIDS[@]}" 2>/dev/null || true + wait "${PIDS[@]}" 2>/dev/null || true + fi +} + +trap cleanup EXIT +trap 'exit 130' INT +trap 'exit 143' TERM + +start_prefill() { + local name="$1" + local devices="$2" + local http_port="$3" + local nccl_port="$4" + + echo "Starting ${name}: GPUs=${devices}, port=${http_port}" + + CUDA_VISIBLE_DEVICES="${devices}" \ + python -m lightllm.server.api_server \ + "${P_COMMON_ARGS[@]}" \ + --run_mode prefill \ + --enable_cpu_cache \ + --cpu_cache_storage_size "${CPU_CACHE_SIZE}" \ + --tp 4 \ + --dp 1 \ + --visual_dp 4 \ + --host 0.0.0.0 \ + --port "${http_port}" \ + --nccl_port "${nccl_port}" \ + >>"${LOG_DIR}/${name}.log" 2>&1 & + + PIDS+=("$!") +} + +start_decode() { + local name="$1" + local devices="$2" + local http_port="$3" + local nccl_port="$4" + + echo "Starting ${name}: GPUs=${devices}, port=${http_port}" + + CUDA_VISIBLE_DEVICES="${devices}" \ + python -m lightllm.server.api_server \ + "${D_COMMON_ARGS[@]}" \ + --run_mode decode \ + --tp 4 \ + --dp 1 \ + --visual_dp 4 \ + --host 0.0.0.0 \ + --port "${http_port}" \ + --nccl_port "${nccl_port}" \ + >>"${LOG_DIR}/${name}.log" 2>&1 & + + PIDS+=("$!") +} + +# Single node: Prefill on GPUs 0-3, Decode on GPUs 4-7. +start_prefill p1 "0,1,2,3" 28761 29761 +start_decode d1 "4,5,6,7" 28762 29762 + +# CPU-only PD Master. This is the public service endpoint. +python -m lightllm.server.api_server \ + --model_dir "${MODEL_DIR}" \ + --model_name "${MODEL_NAME}" \ + --run_mode pd_master \ + --pd_master_mode 1p1d \ + --host 0.0.0.0 \ + --port "${PORT}" \ + --max_image_token_count 4096 \ + --max_image_pixels 3686400 \ + --chat_template "${CHAT_TEMPLATE}" \ + >>"${LOG_DIR}/pd-master.log" 2>&1 & +PIDS+=("$!") + +echo "PD Master is starting at http://${PD_MASTER_IP}:${PORT}" +echo "logs: ${LOG_DIR}" + +# A fixed 1P1D deployment is incomplete as soon as any component exits. +if wait -n "${PIDS[@]}"; then + echo "A LightLLM component exited unexpectedly" >&2 +else + echo "A LightLLM component failed" >&2 +fi +exit 1 From 15641f321bc01ece52a9a7133abc98003a3a7e77 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 21 Aug 2026 13:53:59 +0000 Subject: [PATCH 087/103] fix: pad empty speculative hidden states --- lightllm/utils/custom_kernel_utis.py | 9 +++++++++ unit_tests/common/basemodel/test_model_input.py | 2 ++ unit_tests/utils/test_custom_kernel_utils.py | 12 +++++++++++- 3 files changed, 22 insertions(+), 1 deletion(-) diff --git a/lightllm/utils/custom_kernel_utis.py b/lightllm/utils/custom_kernel_utis.py index 9a7578a243..1a0b46ca0d 100644 --- a/lightllm/utils/custom_kernel_utis.py +++ b/lightllm/utils/custom_kernel_utis.py @@ -121,6 +121,15 @@ def pad2dim_tensor_to_new_batch(input: torch.Tensor, new_batch_size: int): assert input.ndim == 2 origin_batch_size = input.shape[0] hidden = input.shape[1] + + if origin_batch_size == 0: + return torch.zeros( + (new_batch_size, hidden), + dtype=input.dtype, + device=input.device, + requires_grad=False, + ) + out = torch.empty((new_batch_size, hidden), dtype=input.dtype, device=input.device, requires_grad=False) out[0:origin_batch_size, :] = input out[origin_batch_size:, :] = input[0:1, :] diff --git a/unit_tests/common/basemodel/test_model_input.py b/unit_tests/common/basemodel/test_model_input.py index 9853d78cfa..056513bdde 100644 --- a/unit_tests/common/basemodel/test_model_input.py +++ b/unit_tests/common/basemodel/test_model_input.py @@ -158,6 +158,7 @@ def test_padded_prefill_builds_internal_request_for_empty_input(): is_prefill=True, b_prefill_has_output_cpu=[], multimodal_params=[], + mtp_draft_input_hiddens=torch.empty((0, 4), dtype=torch.float32), ) model = SimpleNamespace( mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), @@ -177,6 +178,7 @@ def test_padded_prefill_builds_internal_request_for_empty_input(): assert padded_input.b_req_idx.tolist() == [88] assert padded_input.b_seq_len.tolist() == [1] assert padded_input.b_prefill_has_output_cpu == [False] + assert torch.equal(padded_input.mtp_draft_input_hiddens, torch.zeros((1, 4))) def test_padded_decode_builds_internal_request_from_empty_token_tensor(): diff --git a/unit_tests/utils/test_custom_kernel_utils.py b/unit_tests/utils/test_custom_kernel_utils.py index 47093c9891..4f9d63388d 100644 --- a/unit_tests/utils/test_custom_kernel_utils.py +++ b/unit_tests/utils/test_custom_kernel_utils.py @@ -1,7 +1,7 @@ import torch import time import pytest -from lightllm.utils.custom_kernel_utis import torch_cat_3 +from lightllm.utils.custom_kernel_utis import pad2dim_tensor_to_new_batch, torch_cat_3 def test_torch_cat(): @@ -21,5 +21,15 @@ def test_torch_cat(): return +def test_pad2dim_tensor_to_new_batch(): + input_tensor = torch.tensor([[1.0, 2.0]]) + padded_tensor = pad2dim_tensor_to_new_batch(input=input_tensor, new_batch_size=3) + assert torch.equal(padded_tensor, torch.tensor([[1.0, 2.0], [1.0, 2.0], [1.0, 2.0]])) + + empty_input = torch.empty((0, 2)) + padded_empty_input = pad2dim_tensor_to_new_batch(input=empty_input, new_batch_size=2) + assert torch.equal(padded_empty_input, torch.zeros((2, 2))) + + if __name__ == "__main__": pytest.main() From 8f921287ff8587826bc10e15373aac648d747f65 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Fri, 21 Aug 2026 14:25:05 +0000 Subject: [PATCH 088/103] refactor: encapsulate dp overlap proposal inputs --- .../mode_backend/dp_backend/impl.py | 231 +++++------------- .../mtp_speculative/dp_overlap_engine.py | 21 +- .../dp_overlap_proposers/base.py | 27 +- .../dp_overlap_proposers/eagle3.py | 32 ++- .../dp_overlap_proposers/eagle_no_att.py | 32 ++- .../dp_overlap_proposers/eagle_utils.py | 100 +++++++- .../dp_overlap_proposers/eagle_with_att.py | 32 ++- .../dp_overlap_proposers/vanilla_no_att.py | 31 ++- .../dp_overlap_proposers/vanilla_with_att.py | 30 ++- .../test_dp_overlap_spec_engine.py | 198 ++++++--------- .../mtp_speculative/test_eagle_overlap.py | 61 ++++- .../mtp_speculative/test_vanilla_overlap.py | 50 +++- 12 files changed, 436 insertions(+), 409 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 2ee2f550c4..d6632b86ce 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -35,22 +35,12 @@ def __init__(self) -> None: if spec_mode is not None: if spec_mode in ("dspark", "dflash"): raise NotImplementedError("DP backend does not support DFlash/DSpark parallel block drafting yet.") - self.uses_autoregressive_drafter = spec_mode in ( - "eagle_with_att", - "eagle_no_att", - "eagle3", - ) if self.enable_prefill_microbatch_overlap: self.prefill = self.prefill_overlap_mtp else: self.prefill = self.prefill_mtp if self.enable_decode_microbatch_overlap: self.decode = self.decode_overlap_mtp - self._draft_decode_overlap_func = ( - self._draft_decode_eagle_overlap - if self.uses_autoregressive_drafter - else self._draft_decode_vanilla_overlap - ) else: self.decode = self.decode_mtp else: @@ -459,7 +449,6 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] b_req_idx = model_input.b_req_idx[0:req_num] b_mtp_index = model_input.b_mtp_index[0:req_num] - next_token_ids = torch.empty((0,), dtype=torch.int64, device=model_input.b_req_idx.device) if req_num > 0: ( next_token_ids, @@ -475,6 +464,8 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] b_prefill_has_output_cpu=b_has_out_cpu, mask_func=None, ) + else: + next_token_ids = torch.empty((0,), dtype=torch.int64, device=model_input.b_req_idx.device) # BaseModel 已负责空 batch 的内部 padding,这里直接把真实 target # 输出交给与非 DP 路径相同的 SpecEngine。 @@ -541,17 +532,6 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): selected_rows = async_selected_row_mask_cpu.tensor.tolist() run_reqs = [req for req, selected in zip(run_reqs, selected_rows) if selected] - next_token_ids = torch.empty( - (0,), - dtype=torch.int64, - device=model_input.b_req_idx.device, - ) - b_req_mtp_start_loc = torch.empty( - (0,), - dtype=torch.int32, - device=model_input.b_req_idx.device, - ) - mtp_accept_len = torch.empty_like(b_req_mtp_start_loc) if req_num > 0: next_token_ids, next_token_logprobs = sample( model_output.logits, @@ -585,6 +565,18 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): next_token_ids=next_token_ids, mask=accepted_index == 1, ) + else: + next_token_ids = torch.empty( + (0,), + dtype=torch.int64, + device=model_input.b_req_idx.device, + ) + b_req_mtp_start_loc = torch.empty( + (0,), + dtype=torch.int32, + device=model_input.b_req_idx.device, + ) + mtp_accept_len = torch.empty_like(b_req_mtp_start_loc) verify_event = torch.cuda.Event() verify_event.record() @@ -695,6 +687,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I logits0 = model_output0.logits logits1 = model_output1.logits req_num0, req_num1 = len(run_reqs0), len(run_reqs1) + req_num = req_num0 + req_num1 logits = torch.empty((req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1[0:req_num1, :], non_blocking=True) @@ -706,8 +699,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I b_mtp_index = torch.cat((model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0) b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) - next_token_ids = torch.empty((0,), dtype=torch.int64, device=logits.device) - if (req_num0 + req_num1) > 0: + if req_num > 0: ( next_token_ids, next_token_ids_cpu, @@ -721,6 +713,8 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I is_prefill=True, b_prefill_has_output_cpu=b_has_out_cpu, ) + else: + next_token_ids = torch.empty((0,), dtype=torch.int64, device=logits.device) target_next_token_ids_gpu0 = self._build_padded_next_token_ids( token_ids=next_token_ids, @@ -746,14 +740,14 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I target_next_token_ids1=target_next_token_ids_gpu1, ) - if req_num0 + req_num1 > 0 and g_infer_context.is_linear_att_mixed_model: + if req_num > 0 and g_infer_context.is_linear_att_mixed_model: _b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=_b_req_idx, reqs=run_reqs) sync_event = torch.cuda.Event() sync_event.record() - if req_num0 + req_num1 > 0: + if req_num > 0: event_pack.notify_post_handle_and_wait_pre_post_handle() update_packs = self._pre_post_handle(run_reqs, is_chuncked_mode=not self.disable_chunked_prefill) @@ -784,6 +778,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf run_reqs1, ) = overlap_prepare_decode_inputs(req_objs=decode_reqs) req_num0, req_num1 = len(run_reqs0), len(run_reqs1) + req_num = req_num0 + req_num1 b_mtp_index_cpu0 = model_input0.b_mtp_index b_mtp_index_cpu1 = model_input1.b_mtp_index with torch.cuda.stream(g_infer_context.get_overlap_stream()): @@ -791,11 +786,8 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf logits0 = model_output0.logits logits1 = model_output1.logits run_reqs = run_reqs0 + run_reqs1 - b_req_idx, mtp_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None, None - if (req_num0 + req_num1) > 0: - logits = torch.empty( - (req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device - ) + if req_num > 0: + logits = torch.empty((req_num, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1[0:req_num1, :], non_blocking=True) next_token_ids, next_token_logprobs = sample(logits, run_reqs, self.eos_id) @@ -840,23 +832,35 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf key="mtp_accept_len", gpu_tensor=mtp_accept_len, ) + else: + b_req_idx = torch.empty((0,), dtype=torch.int32, device=model_input0.b_req_idx.device) + mtp_accept_len = torch.empty((0,), dtype=torch.int32, device=model_input0.b_req_idx.device) + b_req_mtp_start_loc = torch.empty((0,), dtype=torch.int32, device=model_input0.b_req_idx.device) + next_token_ids = torch.empty((0,), dtype=torch.int64, device=model_input0.b_req_idx.device) verify_event = torch.cuda.Event() verify_event.record() - extra_mem_indexes_cpu = self._draft_decode_overlap_func( - model_input0=model_input0, - model_input1=model_input1, - model_output0=model_output0, - model_output1=model_output1, - b_req_idx=b_req_idx, - next_token_ids=next_token_ids, - mtp_accept_len=mtp_accept_len, - b_req_mtp_start_loc=b_req_mtp_start_loc, - req_num0=req_num0, - req_num1=req_num1, + proposal = self.decode_draft_engine.propose_next_overlap( + target_model_input0=model_input0, + target_model_output0=model_output0, + target_model_input1=model_input1, + target_model_output1=model_output1, + target_next_token_ids=next_token_ids, + real_verify_rows0=req_num0, + real_verify_rows1=req_num1, + accept_len=mtp_accept_len, ) + if req_num > 0: + mtp_utils.scatter_mtp_next_tokens( + backend=self, + proposal=proposal, + target_next_token_ids=next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + b_req_idx=b_req_idx, + mtp_accept_len=mtp_accept_len, + ) - if (req_num0 + req_num1) > 0: + if req_num > 0: g_infer_context.req_sampling_manager.update_reqs_out_token_counter_gpu( b_req_idx=b_req_idx, next_token_ids=next_token_ids, @@ -865,7 +869,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf sync_event = torch.cuda.Event() sync_event.record() - if req_num0 + req_num1 > 0: + if req_num > 0: event_pack.notify_post_handle_and_wait_pre_post_handle() verify_event.synchronize() mtp_utils.record_request_mtp_metrics( @@ -882,7 +886,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf mem_indexes_cpu = torch.cat( (model_input0.mem_indexes_cpu[0:req_num0], model_input1.mem_indexes_cpu[0:req_num1]), dim=0 ) - extra_mem_indexes_cpu.append( + proposal.extra_mem_indexes_cpu.append( MtpMemIndexesToFree( mem_indexes_cpu=mem_indexes_cpu, free_mask_cpu=accepted_index_cpu == 0, @@ -900,141 +904,16 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf ) mtp_utils.free_mem_indexes( backend=self, - extra_mem_indexes_cpu=extra_mem_indexes_cpu, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, ) event_pack.notify_pre_post_handle() else: event_pack.notify_post_handle_and_wait_pre_post_handle() event_pack.notify_forward_and_wait_post_handle() - event_pack.notify_pre_post_handle() - return - - def _draft_decode_vanilla_overlap( - self, - model_input0: ModelInput, - model_input1: ModelInput, - model_output0: ModelOutput, - model_output1: ModelOutput, - b_req_idx: torch.Tensor, - next_token_ids: torch.Tensor = None, - mtp_accept_len: torch.Tensor = None, - b_req_mtp_start_loc: torch.Tensor = None, - req_num0: int = 0, - req_num1: int = 0, - ): - if mtp_accept_len is None: - mtp_accept_len = torch.empty((0,), dtype=torch.int32, device=model_input0.b_req_idx.device) - verify_width = self.max_draft_step + 1 - real_request_num0 = req_num0 // verify_width - padded_next_token_ids0 = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input0.batch_size, - copy_len=req_num0, - device=model_input0.b_req_idx.device, - source_start=0, - ) - padded_next_token_ids1 = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input1.batch_size, - copy_len=req_num1, - device=model_input1.b_req_idx.device, - source_start=req_num0, - ) - - proposal = self.decode_draft_engine.propose_next_overlap( - target_model_input0=model_input0, - target_model_output0=model_output0, - target_next_token_ids0=padded_next_token_ids0, - real_verify_rows0=req_num0, - accept_len0=mtp_accept_len[:real_request_num0], - target_model_input1=model_input1, - target_model_output1=model_output1, - target_next_token_ids1=padded_next_token_ids1, - real_verify_rows1=req_num1, - accept_len1=mtp_accept_len[real_request_num0:], - ) - - if req_num0 + req_num1 > 0: - mtp_utils.scatter_mtp_next_tokens( - backend=self, - proposal=proposal, - target_next_token_ids=next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, - ) - return proposal.extra_mem_indexes_cpu - - def _draft_decode_eagle_overlap( - self, - model_input0: ModelInput, - model_input1: ModelInput, - model_output0: ModelOutput, - model_output1: ModelOutput, - b_req_idx: torch.Tensor, - next_token_ids: torch.Tensor = None, - mtp_accept_len: torch.Tensor = None, - b_req_mtp_start_loc: torch.Tensor = None, - req_num0: int = 0, - req_num1: int = 0, - ): - verify_width = self.max_draft_step + 1 - real_request_num0 = req_num0 // verify_width - real_request_num1 = req_num1 // verify_width - request_capacity0 = model_input0.batch_size // verify_width - request_capacity1 = model_input1.batch_size // verify_width - - padded_next_token_ids0 = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input0.batch_size, - copy_len=req_num0, - device=model_input0.b_req_idx.device, - source_start=0, - ) - padded_next_token_ids1 = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input1.batch_size, - copy_len=req_num1, - device=model_input1.b_req_idx.device, - source_start=req_num0, - ) - padded_accept_len0 = torch.ones( - (request_capacity0,), - dtype=torch.int32, - device=model_input0.b_req_idx.device, - ) - padded_accept_len1 = torch.ones( - (request_capacity1,), - dtype=torch.int32, - device=model_input1.b_req_idx.device, - ) - if real_request_num0 > 0: - padded_accept_len0[:real_request_num0].copy_(mtp_accept_len[:real_request_num0]) - if real_request_num1 > 0: - padded_accept_len1[:real_request_num1].copy_( - mtp_accept_len[real_request_num0 : real_request_num0 + real_request_num1] - ) - - proposal = self.decode_draft_engine.propose_next_overlap( - target_model_input0=model_input0, - target_model_output0=model_output0, - target_next_token_ids0=padded_next_token_ids0, - real_verify_rows0=req_num0, - accept_len0=padded_accept_len0, - target_model_input1=model_input1, - target_model_output1=model_output1, - target_next_token_ids1=padded_next_token_ids1, - real_verify_rows1=req_num1, - accept_len1=padded_accept_len1, - ) - - if req_num0 + req_num1 > 0: - mtp_utils.scatter_mtp_next_tokens( + sync_event.synchronize() + mtp_utils.free_mem_indexes( backend=self, - proposal=proposal, - target_next_token_ids=next_token_ids, - b_req_mtp_start_loc=b_req_mtp_start_loc, - b_req_idx=b_req_idx, - mtp_accept_len=mtp_accept_len, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, ) - return proposal.extra_mem_indexes_cpu + event_pack.notify_pre_post_handle() + return diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index 0921422fc2..8bae91971e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING import torch @@ -50,26 +50,25 @@ def propose_next_overlap( self, target_model_input0: ModelInput, # batch_size = verify_batch_size0 target_model_output0: ModelOutput, # logits: [verify_batch_size0, vocab_size] - target_next_token_ids0: torch.Tensor, # [verify_batch_size0] - real_verify_rows0: int, - accept_len0: Optional[torch.Tensor], # [req_num0] target_model_input1: ModelInput, # batch_size = verify_batch_size1 target_model_output1: ModelOutput, # logits: [verify_batch_size1, vocab_size] - target_next_token_ids1: torch.Tensor, # [verify_batch_size1] + target_next_token_ids: torch.Tensor, # [real_verify_rows0 + real_verify_rows1] + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: Optional[torch.Tensor], # [req_num1] + accept_len: torch.Tensor, # [real_req_num0 + real_req_num1] ) -> SpecProposal: + real_verify_rows = real_verify_rows0 + real_verify_rows1 + assert target_next_token_ids.shape == (real_verify_rows,) + return self.proposer.propose_next_overlap( target_model_input0=target_model_input0, target_model_output0=target_model_output0, - target_next_token_ids0=target_next_token_ids0, - real_verify_rows0=real_verify_rows0, - accept_len0=accept_len0, target_model_input1=target_model_input1, target_model_output1=target_model_output1, - target_next_token_ids1=target_next_token_ids1, + target_next_token_ids=target_next_token_ids, + real_verify_rows0=real_verify_rows0, real_verify_rows1=real_verify_rows1, - accept_len1=accept_len1, + accept_len=accept_len, draft_step=self.planner.get_draft_step(), ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py index 7e3877894c..2cc880b647 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -36,16 +36,33 @@ def propose_next_overlap( self, target_model_input0: ModelInput, # batch_size = padded_verify_batch_size0 target_model_output0: ModelOutput, # logits: [padded_verify_batch_size0, vocab_size] - target_next_token_ids0: torch.Tensor, # [padded_verify_batch_size0] - real_verify_rows0: int, - accept_len0: torch.Tensor | None, # [req_num0] target_model_input1: ModelInput, # batch_size = padded_verify_batch_size1 target_model_output1: ModelOutput, # logits: [padded_verify_batch_size1, vocab_size] - target_next_token_ids1: torch.Tensor, # [padded_verify_batch_size1] + target_next_token_ids: torch.Tensor, # [real_verify_rows0 + real_verify_rows1] + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: torch.Tensor | None, # [req_num1] + accept_len: torch.Tensor, # [real_req_num0 + real_req_num1] draft_step: int, ) -> SpecProposal: """Generate one proposal from two DP-overlapped decode microbatches.""" raise NotImplementedError + + @staticmethod + def _build_padded_next_token_ids( + target_next_token_ids: torch.Tensor, + batch_size: int, + real_verify_rows: int, + device: torch.device, + source_start: int, + ) -> torch.Tensor: + """将合并后的真实 target token 拆分并补齐到一个 microbatch 的 verify shape。""" + + assert 0 <= real_verify_rows <= batch_size + padded_token_ids = torch.zeros((batch_size,), dtype=torch.int64, device=device) + if real_verify_rows > 0: + padded_token_ids[:real_verify_rows].copy_( + target_next_token_ids[source_start : source_start + real_verify_rows], + non_blocking=True, + ) + return padded_token_ids diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index 69fc25928e..e5a7e287c8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -38,28 +38,24 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor | None, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, + target_next_token_ids: torch.Tensor, + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: torch.Tensor | None, + accept_len: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: return propose_next_dp_eagle_autoregressive_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - real_verify_rows0, - accept_len0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - real_verify_rows1, - accept_len1, - draft_step, - self._map_draft_token_ids, + proposer=self, + target_model_input0=target_model_input0, + target_model_output0=target_model_output0, + target_model_input1=target_model_input1, + target_model_output1=target_model_output1, + target_next_token_ids=target_next_token_ids, + real_verify_rows0=real_verify_rows0, + real_verify_rows1=real_verify_rows1, + accept_len=accept_len, + draft_step=draft_step, + map_draft_token_ids=self._map_draft_token_ids, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index fe151f6a83..c2250d5d51 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -35,28 +35,24 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor | None, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, + target_next_token_ids: torch.Tensor, + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: torch.Tensor | None, + accept_len: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: return propose_next_dp_eagle_fixed_layout_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - real_verify_rows0, - accept_len0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - real_verify_rows1, - accept_len1, - draft_step, - lambda token_ids: token_ids, + proposer=self, + target_model_input0=target_model_input0, + target_model_output0=target_model_output0, + target_model_input1=target_model_input1, + target_model_output1=target_model_output1, + target_next_token_ids=target_next_token_ids, + real_verify_rows0=real_verify_rows0, + real_verify_rows1=real_verify_rows1, + accept_len=accept_len, + draft_step=draft_step, + map_draft_token_ids=lambda token_ids: token_ids, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index 72f3fcac08..de81db3e6f 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -96,23 +96,90 @@ def pad_dp_step_mem_indexes( return padded +def prepare_dp_eagle_overlap_decode_inputs( + proposer: BaseDpOverlapProposer, + target_model_input0: ModelInput, + target_model_input1: ModelInput, + target_next_token_ids: torch.Tensor, + real_verify_rows0: int, + real_verify_rows1: int, + accept_len: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """拆分真实 verify 结果,并补齐为两个 EAGLE overlap microbatch 的固定布局。""" + + verify_width = proposer.backend.max_draft_step + 1 + real_request_num0 = real_verify_rows0 // verify_width + real_request_num1 = real_verify_rows1 // verify_width + request_capacity0 = target_model_input0.batch_size // verify_width + request_capacity1 = target_model_input1.batch_size // verify_width + + assert accept_len.shape == (real_request_num0 + real_request_num1,) + + target_next_token_ids0 = proposer._build_padded_next_token_ids( + target_next_token_ids=target_next_token_ids, + batch_size=target_model_input0.batch_size, + real_verify_rows=real_verify_rows0, + device=target_model_input0.b_req_idx.device, + source_start=0, + ) + target_next_token_ids1 = proposer._build_padded_next_token_ids( + target_next_token_ids=target_next_token_ids, + batch_size=target_model_input1.batch_size, + real_verify_rows=real_verify_rows1, + device=target_model_input1.b_req_idx.device, + source_start=real_verify_rows0, + ) + + padded_accept_len0 = torch.ones( + (request_capacity0,), + dtype=torch.int32, + device=target_model_input0.b_req_idx.device, + ) + padded_accept_len1 = torch.ones( + (request_capacity1,), + dtype=torch.int32, + device=target_model_input1.b_req_idx.device, + ) + if real_request_num0 > 0: + padded_accept_len0[:real_request_num0].copy_(accept_len[:real_request_num0]) + if real_request_num1 > 0: + padded_accept_len1[:real_request_num1].copy_( + accept_len[real_request_num0 : real_request_num0 + real_request_num1] + ) + + return target_next_token_ids0, target_next_token_ids1, padded_accept_len0, padded_accept_len1 + + def propose_next_dp_eagle_autoregressive_overlap( proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor | None, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, + target_next_token_ids: torch.Tensor, + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: torch.Tensor | None, + accept_len: torch.Tensor, draft_step: int, map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], ) -> EagleSpecProposal: """运行 DP EAGLE extend 后接单 token overlap decode 的 proposal 流程。""" + ( + target_next_token_ids0, + target_next_token_ids1, + accept_len0, + accept_len1, + ) = prepare_dp_eagle_overlap_decode_inputs( + proposer=proposer, + target_model_input0=target_model_input0, + target_model_input1=target_model_input1, + target_next_token_ids=target_next_token_ids, + real_verify_rows0=real_verify_rows0, + real_verify_rows1=real_verify_rows1, + accept_len=accept_len, + ) + verify_width = proposer.backend.max_draft_step + 1 model_inputs = (target_model_input0, target_model_input1) model_outputs = (target_model_output0, target_model_output1) @@ -250,19 +317,32 @@ def propose_next_dp_eagle_fixed_layout_overlap( proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, + target_next_token_ids: torch.Tensor, + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: torch.Tensor, + accept_len: torch.Tensor, draft_step: int, map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], ) -> EagleSpecProposal: """以 expanded verify-row layout 运行 decode,返回按真实请求压缩的 proposal。""" + ( + target_next_token_ids0, + target_next_token_ids1, + accept_len0, + accept_len1, + ) = prepare_dp_eagle_overlap_decode_inputs( + proposer=proposer, + target_model_input0=target_model_input0, + target_model_input1=target_model_input1, + target_next_token_ids=target_next_token_ids, + real_verify_rows0=real_verify_rows0, + real_verify_rows1=real_verify_rows1, + accept_len=accept_len, + ) + verify_width = proposer.backend.max_draft_step + 1 model_inputs = (target_model_input0, target_model_input1) real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index c470a72445..dc0e1609f3 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -35,28 +35,24 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor | None, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, + target_next_token_ids: torch.Tensor, + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: torch.Tensor | None, + accept_len: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: return propose_next_dp_eagle_fixed_layout_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - real_verify_rows0, - accept_len0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - real_verify_rows1, - accept_len1, - draft_step, - lambda token_ids: token_ids, + proposer=self, + target_model_input0=target_model_input0, + target_model_output0=target_model_output0, + target_model_input1=target_model_input1, + target_model_output1=target_model_output1, + target_next_token_ids=target_next_token_ids, + real_verify_rows0=real_verify_rows0, + real_verify_rows1=real_verify_rows1, + accept_len=accept_len, + draft_step=draft_step, + map_draft_token_ids=lambda token_ids: token_ids, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index e29f805430..81ca39de3d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -25,29 +25,40 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor | None, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, + target_next_token_ids: torch.Tensor, + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: torch.Tensor | None, + accept_len: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: - assert accept_len0 is not None - assert accept_len1 is not None - verify_width = self.backend.max_draft_step + 1 real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) req_num_by_batch = tuple(row_count // verify_width for row_count in real_verify_rows) + assert accept_len.shape == (sum(req_num_by_batch),) + + target_next_token_ids0 = self._build_padded_next_token_ids( + target_next_token_ids=target_next_token_ids, + batch_size=target_model_input0.batch_size, + real_verify_rows=real_verify_rows0, + device=target_model_input0.b_req_idx.device, + source_start=0, + ) + target_next_token_ids1 = self._build_padded_next_token_ids( + target_next_token_ids=target_next_token_ids, + batch_size=target_model_input1.batch_size, + real_verify_rows=real_verify_rows1, + device=target_model_input1.b_req_idx.device, + source_start=real_verify_rows0, + ) req_start_rows = ( torch.arange(0, real_verify_rows0, verify_width, device=target_next_token_ids0.device), torch.arange(0, real_verify_rows1, verify_width, device=target_next_token_ids1.device), ) accepted_tail_rows = ( - req_start_rows[0] + accept_len0[: req_num_by_batch[0]] - 1, - req_start_rows[1] + accept_len1[: req_num_by_batch[1]] - 1, + req_start_rows[0] + accept_len[: req_num_by_batch[0]] - 1, + req_start_rows[1] + accept_len[req_num_by_batch[0] : sum(req_num_by_batch)] - 1, ) model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] draft_token_ids = [target_next_token_ids0, target_next_token_ids1] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index c1f8e3d6de..895c55dbdf 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -52,31 +52,43 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - real_verify_rows0: int, - accept_len0: torch.Tensor | None, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, + target_next_token_ids: torch.Tensor, + real_verify_rows0: int, real_verify_rows1: int, - accept_len1: torch.Tensor | None, + accept_len: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: - assert accept_len0 is not None - assert accept_len1 is not None assert draft_step == self.backend.max_draft_step assert len(self.backend.draft_models) == draft_step verify_width = self.backend.max_draft_step + 1 real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) req_num_by_batch = tuple(row_count // verify_width for row_count in real_verify_rows) + assert accept_len.shape == (sum(req_num_by_batch),) + + target_next_token_ids0 = self._build_padded_next_token_ids( + target_next_token_ids=target_next_token_ids, + batch_size=target_model_input0.batch_size, + real_verify_rows=real_verify_rows0, + device=target_model_input0.b_req_idx.device, + source_start=0, + ) + target_next_token_ids1 = self._build_padded_next_token_ids( + target_next_token_ids=target_next_token_ids, + batch_size=target_model_input1.batch_size, + real_verify_rows=real_verify_rows1, + device=target_model_input1.b_req_idx.device, + source_start=real_verify_rows0, + ) req_start_rows = ( torch.arange(0, real_verify_rows0, verify_width, device=target_next_token_ids0.device), torch.arange(0, real_verify_rows1, verify_width, device=target_next_token_ids1.device), ) accept_len_by_batch = ( - accept_len0[: req_num_by_batch[0]], - accept_len1[: req_num_by_batch[1]], + accept_len[: req_num_by_batch[0]], + accept_len[req_num_by_batch[0] : sum(req_num_by_batch)], ) accepted_tail_rows = tuple(starts + lengths - 1 for starts, lengths in zip(req_start_rows, accept_len_by_batch)) model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 1ede7c4acb..03bf193f68 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -16,47 +16,7 @@ LightSpecPlanner, SpecDecodePlan, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( - MtpMemIndexesToFree, - SpecProposal, -) - - -class _RecordingSpecEngine: - def __init__(self): - self.propose_args = None - self.propose_overlap_args = None - - def propose_next(self, **kwargs): - self.propose_args = kwargs - token_ids = kwargs["target_next_token_ids"].new_zeros((kwargs["b_req_mtp_start_loc"].shape[0], 7)) - return SpecProposal( - token_ids=token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([123], dtype=torch.int32))], - ) - - def propose_next_overlap(self, **kwargs): - self.propose_overlap_args = kwargs - request_count = (kwargs["real_verify_rows0"] + kwargs["real_verify_rows1"]) // 8 - token_ids = kwargs["target_next_token_ids0"].new_zeros((request_count, 7)) - return SpecProposal( - token_ids=token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=torch.tensor([456], dtype=torch.int32))], - ) - - -def _capture_scatter_args(monkeypatch): - scatter_args = {} - monkeypatch.setattr( - dp_backend_impl.mtp_utils, "scatter_mtp_next_tokens", lambda **kwargs: scatter_args.update(kwargs) - ) - return scatter_args - - -def _assert_all_mem_indexes_are_freed(extra_mem_indexes_cpu, expected): - assert len(extra_mem_indexes_cpu) == 1 - assert torch.equal(extra_mem_indexes_cpu[0].mem_indexes_cpu, expected) - assert extra_mem_indexes_cpu[0].free_mask_cpu is None +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal def test_dp_backend_reuses_common_engine_outside_overlap(): @@ -315,90 +275,94 @@ def propose_next(self, **kwargs): assert calls == ["plan", "prepare", "propose", "post_wait", "forward_wait", "free", "pre_post"] -def test_dp_overlap_eagle_passes_both_fixed_verify_layouts_to_proposer(monkeypatch): - scatter_args = _capture_scatter_args(monkeypatch) - backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) - backend.max_draft_step = 7 - backend.spec_engine = _RecordingSpecEngine() - backend.dp_overlap_spec_engine = _RecordingSpecEngine() - backend.decode_draft_engine = backend.dp_overlap_spec_engine - model_input0 = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), - ) - model_input1 = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), - ) - model_output0 = SimpleNamespace(spec_hidden=torch.randn(16, 4)) - model_output1 = SimpleNamespace(spec_hidden=torch.randn(16, 4)) - next_token_ids = torch.arange(24, dtype=torch.int64) - b_req_idx = torch.arange(24, dtype=torch.int32) - start_locs = torch.tensor([0, 8, 16], dtype=torch.int32) +def test_dp_overlap_engine_delegates_raw_verify_layout_to_proposer(): + calls = {} + + class _Proposer: + def propose_next_overlap(self, **kwargs): + calls.update(kwargs) + return SpecProposal(token_ids=kwargs["target_next_token_ids"].new_empty((3, 7))) + + engine = DPOverlapSpecEngine.__new__(DPOverlapSpecEngine) + engine.proposer = _Proposer() + engine.planner = SimpleNamespace(get_draft_step=lambda: 7) + model_input0 = SimpleNamespace(batch_size=16) + model_input1 = SimpleNamespace(batch_size=16) + model_output0 = SimpleNamespace() + model_output1 = SimpleNamespace() + target_next_token_ids = torch.arange(24, dtype=torch.int64) accept_len = torch.tensor([2, 3, 4], dtype=torch.int32) - extra_mem = backend._draft_decode_eagle_overlap( - model_input0=model_input0, - model_output0=model_output0, - model_input1=model_input1, - model_output1=model_output1, - b_req_idx=b_req_idx, - next_token_ids=next_token_ids, - mtp_accept_len=accept_len, - b_req_mtp_start_loc=start_locs, - req_num0=8, - req_num1=16, + proposal = engine.propose_next_overlap( + target_model_input0=model_input0, + target_model_output0=model_output0, + target_model_input1=model_input1, + target_model_output1=model_output1, + target_next_token_ids=target_next_token_ids, + real_verify_rows0=8, + real_verify_rows1=16, + accept_len=accept_len, ) - propose_args = backend.dp_overlap_spec_engine.propose_overlap_args - assert propose_args["target_next_token_ids0"].shape == (16,) - assert propose_args["target_next_token_ids1"].shape == (16,) - assert propose_args["real_verify_rows0"] == 8 - assert propose_args["real_verify_rows1"] == 16 - assert torch.equal(propose_args["accept_len0"], torch.tensor([2, 1], dtype=torch.int32)) - assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) - assert scatter_args["proposal"].token_ids.shape == (3, 7) - assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) - _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([456], dtype=torch.int32)) + assert calls["target_model_input0"] is model_input0 + assert calls["target_model_input1"] is model_input1 + assert calls["target_model_output0"] is model_output0 + assert calls["target_model_output1"] is model_output1 + assert calls["target_next_token_ids"] is target_next_token_ids + assert calls["real_verify_rows0"] == 8 + assert calls["real_verify_rows1"] == 16 + assert calls["accept_len"] is accept_len + assert calls["draft_step"] == 7 + assert proposal.token_ids.shape == (3, 7) -def test_dp_overlap_vanilla_delegates_both_microbatches_to_proposer(monkeypatch): - scatter_args = _capture_scatter_args(monkeypatch) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_dp_overlap_decode_delegates_empty_layout_and_frees_proposal(monkeypatch): + device = "cuda" + empty_i32 = torch.empty((0,), dtype=torch.int32, device=device) + model_input0 = SimpleNamespace(batch_size=0, b_req_idx=empty_i32, b_mtp_index=empty_i32) + model_input1 = SimpleNamespace(batch_size=0, b_req_idx=empty_i32, b_mtp_index=empty_i32) + model_output0 = SimpleNamespace(logits=torch.empty((0, 8), device=device)) + model_output1 = SimpleNamespace(logits=torch.empty((0, 8), device=device)) + calls = [] + + class _OverlapEngine: + def propose_next_overlap(self, **kwargs): + calls.append("propose") + assert kwargs["target_next_token_ids"].shape == (0,) + assert kwargs["target_next_token_ids"].dtype == torch.int64 + assert kwargs["real_verify_rows0"] == 0 + assert kwargs["real_verify_rows1"] == 0 + assert kwargs["accept_len"].shape == (0,) + assert kwargs["accept_len"].dtype == torch.int32 + return SpecProposal(token_ids=torch.empty((0, 2), dtype=torch.int64, device=device)) + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) - backend.max_draft_step = 7 - backend.spec_engine = _RecordingSpecEngine() - backend.dp_overlap_spec_engine = _RecordingSpecEngine() - backend.decode_draft_engine = backend.dp_overlap_spec_engine - model_input0 = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), + backend.decode_draft_engine = _OverlapEngine() + backend.model = SimpleNamespace( + microbatch_overlap_decode=lambda input0, input1: (model_output0, model_output1), ) - model_input1 = SimpleNamespace( - batch_size=16, - b_req_idx=torch.arange(16, dtype=torch.int32), + event_pack = SimpleNamespace( + notify_post_handle_and_wait_pre_post_handle=lambda: calls.append("post_wait"), + notify_forward_and_wait_post_handle=lambda: calls.append("forward_wait"), + notify_pre_post_handle=lambda: calls.append("pre_post"), ) - next_token_ids = torch.arange(24, dtype=torch.int64) - - extra_mem = backend._draft_decode_vanilla_overlap( - model_input0=model_input0, - model_output0=SimpleNamespace(), - model_input1=model_input1, - model_output1=SimpleNamespace(), - b_req_idx=torch.arange(24, dtype=torch.int32), - next_token_ids=next_token_ids, - mtp_accept_len=torch.tensor([2, 3, 4], dtype=torch.int32), - b_req_mtp_start_loc=torch.tensor([0, 8, 16], dtype=torch.int32), - req_num0=8, - req_num1=16, + monkeypatch.setattr( + dp_backend_impl, + "overlap_prepare_decode_inputs", + lambda req_objs: (model_input0, [], model_input1, []), ) + monkeypatch.setattr( + dp_backend_impl, + "g_infer_context", + SimpleNamespace(get_overlap_stream=lambda: torch.cuda.current_stream()), + ) + monkeypatch.setattr( + dp_backend_impl.mtp_utils, + "free_mem_indexes", + lambda **kwargs: calls.append("free"), + ) + + backend.decode_overlap_mtp(event_pack=event_pack, decode_reqs=[]) - propose_args = backend.dp_overlap_spec_engine.propose_overlap_args - assert propose_args["target_next_token_ids0"].shape == (16,) - assert propose_args["target_next_token_ids1"].shape == (16,) - assert propose_args["real_verify_rows0"] == 8 - assert propose_args["real_verify_rows1"] == 16 - assert torch.equal(propose_args["accept_len0"], torch.tensor([2], dtype=torch.int32)) - assert torch.equal(propose_args["accept_len1"], torch.tensor([3, 4], dtype=torch.int32)) - assert scatter_args["proposal"].token_ids.shape == (3, 7) - assert torch.equal(scatter_args["target_next_token_ids"], next_token_ids) - _assert_all_mem_indexes_are_freed(extra_mem, torch.tensor([456], dtype=torch.int32)) + assert calls == ["propose", "post_wait", "forward_wait", "free", "pre_post"] diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 44208139dc..88ddc8dcfe 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -90,16 +90,17 @@ def test_overlap_eagle_keeps_fixed_verify_layout(monkeypatch): target_model_output0=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), - target_next_token_ids0=torch.arange(6, dtype=torch.int64), - real_verify_rows0=3, - accept_len0=torch.tensor([2, 1], dtype=torch.int32), target_model_input1=model_input1, target_model_output1=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), - target_next_token_ids1=torch.arange(10, 16, dtype=torch.int64), + target_next_token_ids=torch.cat( + (torch.arange(3, dtype=torch.int64), torch.arange(10, 16, dtype=torch.int64)), + dim=0, + ), + real_verify_rows0=3, real_verify_rows1=6, - accept_len1=torch.tensor([1, 3], dtype=torch.int32), + accept_len=torch.tensor([2, 1, 3], dtype=torch.int32), draft_step=2, ) @@ -116,6 +117,45 @@ def test_overlap_eagle_keeps_fixed_verify_layout(monkeypatch): assert torch.equal(model_input1.mem_indexes, torch.tensor([2, 2, 4, 5, 3, 5], dtype=torch.int32)) +def test_overlap_eagle_builds_padded_inputs_for_empty_verify_rows(monkeypatch): + draft_model = _DraftModel() + backend = SimpleNamespace( + max_draft_step=2, + draft_models=[draft_model], + model=SimpleNamespace( + req_manager=SimpleNamespace( + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + ) + ), + _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), + ) + proposer = DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=False) + monkeypatch.setattr( + mtp_utils, + "alloc_mem_indexes", + lambda token_count: torch.arange(token_count, dtype=torch.int32), + ) + + proposal = proposer.propose_next_overlap( + target_model_input0=_target_input(batch_size=6), + target_model_output0=ModelOutput( + logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + ), + target_model_input1=_target_input(batch_size=6), + target_model_output1=ModelOutput( + logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + ), + target_next_token_ids=torch.empty((0,), dtype=torch.int64), + real_verify_rows0=0, + real_verify_rows1=0, + accept_len=torch.empty((0,), dtype=torch.int32), + draft_step=2, + ) + + assert proposal.token_ids.shape == (0, 2) + assert draft_model.decode_batch_sizes == [(6, 6), (6, 6)] + + def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): draft_model = _DraftModel() backend = SimpleNamespace( @@ -142,16 +182,17 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): target_model_output0=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), - target_next_token_ids0=torch.arange(6, dtype=torch.int64), - real_verify_rows0=3, - accept_len0=torch.tensor([2, 1], dtype=torch.int32), target_model_input1=model_input1, target_model_output1=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) ), - target_next_token_ids1=torch.arange(10, 16, dtype=torch.int64), + target_next_token_ids=torch.cat( + (torch.arange(3, dtype=torch.int64), torch.arange(10, 16, dtype=torch.int64)), + dim=0, + ), + real_verify_rows0=3, real_verify_rows1=6, - accept_len1=torch.tensor([1, 3], dtype=torch.int32), + accept_len=torch.tensor([2, 1, 3], dtype=torch.int32), draft_step=2, ) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py index b6ca60da83..cc17aba811 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -43,8 +43,8 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) - model_input0 = SimpleNamespace(batch_size=6) - model_input1 = SimpleNamespace(batch_size=6) + model_input0 = SimpleNamespace(batch_size=6, b_req_idx=torch.arange(6, dtype=torch.int32, device=device)) + model_input1 = SimpleNamespace(batch_size=6, b_req_idx=torch.arange(6, dtype=torch.int32, device=device)) proposal = proposer.propose_next_overlap( target_model_input0=model_input0, @@ -52,17 +52,15 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): logits=torch.empty((6, 1), device=device), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), ), - target_next_token_ids0=torch.tensor([10, 11, 0, 0, 0, 0], dtype=torch.int64, device=device), - real_verify_rows0=3, - accept_len0=torch.tensor([2], dtype=torch.int32, device=device), target_model_input1=model_input1, target_model_output1=ModelOutput( logits=torch.empty((6, 1), device=device), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), ), - target_next_token_ids1=torch.tensor([20, 21, 22, 0, 0, 0], dtype=torch.int64, device=device), + target_next_token_ids=torch.tensor([10, 11, 0, 20, 21, 22], dtype=torch.int64, device=device), + real_verify_rows0=3, real_verify_rows1=3, - accept_len1=torch.tensor([1], dtype=torch.int32, device=device), + accept_len=torch.tensor([2, 1], dtype=torch.int32, device=device), draft_step=2, ) @@ -73,3 +71,41 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): assert proposal.extra_mem_indexes_cpu == [] assert draft_models[0].decode_batch_sizes == [(6, 6)] assert draft_models[1].decode_batch_sizes == [(6, 6)] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_dp_vanilla_proposer_builds_padded_inputs_for_empty_verify_rows(): + device = "cuda" + draft_models = [_DraftModel(), _DraftModel()] + backend = SimpleNamespace( + max_draft_step=2, + draft_models=draft_models, + _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), + ) + proposer = DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) + model_input0 = SimpleNamespace(batch_size=6, b_req_idx=torch.arange(6, dtype=torch.int32, device=device)) + model_input1 = SimpleNamespace(batch_size=6, b_req_idx=torch.arange(6, dtype=torch.int32, device=device)) + model_output0 = ModelOutput( + logits=torch.empty((6, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), + ) + model_output1 = ModelOutput( + logits=torch.empty((6, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), + ) + + proposal = proposer.propose_next_overlap( + target_model_input0=model_input0, + target_model_output0=model_output0, + target_model_input1=model_input1, + target_model_output1=model_output1, + target_next_token_ids=torch.empty((0,), dtype=torch.int64, device=device), + real_verify_rows0=0, + real_verify_rows1=0, + accept_len=torch.empty((0,), dtype=torch.int32, device=device), + draft_step=2, + ) + + assert proposal.token_ids.shape == (0, 2) + assert draft_models[0].decode_batch_sizes == [(6, 6)] + assert draft_models[1].decode_batch_sizes == [(6, 6)] From acb6532c7c5dcd99d632d2af0b196d3ce1482925 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sat, 22 Aug 2026 03:21:38 +0000 Subject: [PATCH 089/103] refactor dp overlap speculative scheduling --- lightllm/common/basemodel/cuda_graph.py | 3 + .../mode_backend/dp_backend/impl.py | 145 ++++--- .../mtp_speculative/dp_overlap_engine.py | 147 ++++++- .../dp_overlap_planner/__init__.py | 14 - .../dp_overlap_planner/base.py | 11 - .../dp_overlap_planner/fixed.py | 11 - .../dp_overlap_proposers/base.py | 30 +- .../dp_overlap_proposers/eagle3.py | 22 +- .../dp_overlap_proposers/eagle_no_att.py | 27 +- .../dp_overlap_proposers/eagle_utils.py | 399 +++++++++--------- .../dp_overlap_proposers/eagle_with_att.py | 26 +- .../dp_overlap_proposers/vanilla_no_att.py | 88 ++-- .../dp_overlap_proposers/vanilla_with_att.py | 83 ++-- .../mtp_speculative/planner/lightspec.py | 6 +- .../test_dp_overlap_spec_engine.py | 185 ++++++-- .../mtp_speculative/test_eagle_overlap.py | 243 ++++++++--- .../mtp_speculative/test_planner.py | 134 ++++-- .../mtp_speculative/test_vanilla_overlap.py | 76 +++- 18 files changed, 1083 insertions(+), 567 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/base.py delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/fixed.py diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 5fde98eab8..5849cccf54 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -189,6 +189,9 @@ def _measure_replay_cost(self, graph_obj: torch.cuda.CUDAGraph, batch_size: int) device="cuda", ) dist.all_reduce(infer_cost_ms_tensor, op=dist.ReduceOp.MIN, group=dist.group.WORLD) + if self.enable_decode_microbatch_overlap: + # overlap graph 每次 replay 同时处理两个等容量 microbatch。 + batch_size *= 2 self.infer_cost_ms_by_batch_size[batch_size] = float(infer_cost_ms_tensor.item()) def capture_decode( diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index d6632b86ce..de93987c43 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -70,7 +70,10 @@ def init_spec_engine(self): enable_dynmaic_mtp=self.args.mtp_dynamic_verify, ) - self.dp_overlap_spec_engine = DPOverlapSpecEngine(**engine_kwargs) + self.dp_overlap_spec_engine = DPOverlapSpecEngine( + **engine_kwargs, + common_engine=self.spec_engine, + ) self.prefill_draft_engine = ( self.dp_overlap_spec_engine if self.enable_prefill_microbatch_overlap else self.spec_engine ) @@ -308,11 +311,7 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer logits1 = model_output1.logits req_num0, req_num1 = len(run_reqs0), len(run_reqs1) - logits = torch.empty( - (req_num0 + req_num1, logits0.shape[1]), - dtype=logits0.dtype, - device=logits0.device, - ) + logits = torch.empty((req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) logits[req_num0 : req_num0 + req_num1, :].copy_(logits1[0:req_num1, :], non_blocking=True) @@ -320,20 +319,8 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer b_has_out_cpu = ( model_input0.b_prefill_has_output_cpu[0:req_num0] + model_input1.b_prefill_has_output_cpu[0:req_num1] ) - b_mtp_index = torch.cat( - ( - model_input0.b_mtp_index[0:req_num0], - model_input1.b_mtp_index[0:req_num1], - ), - dim=0, - ) - b_req_idx = torch.cat( - ( - model_input0.b_req_idx[0:req_num0], - model_input1.b_req_idx[0:req_num1], - ), - dim=0, - ) + b_mtp_index = torch.cat((model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0) + b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) if req_num0 + req_num1 > 0: ( @@ -688,7 +675,11 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I logits1 = model_output1.logits req_num0, req_num1 = len(run_reqs0), len(run_reqs1) req_num = req_num0 + req_num1 - logits = torch.empty((req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) + logits = torch.empty( + (req_num0 + req_num1, logits0.shape[1]), + dtype=logits0.dtype, + device=logits0.device, + ) logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1[0:req_num1, :], non_blocking=True) @@ -696,8 +687,20 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I b_has_out_cpu = ( model_input0.b_prefill_has_output_cpu[0:req_num0] + model_input1.b_prefill_has_output_cpu[0:req_num1] ) - b_mtp_index = torch.cat((model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0) - b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) + b_mtp_index = torch.cat( + ( + model_input0.b_mtp_index[0:req_num0], + model_input1.b_mtp_index[0:req_num1], + ), + dim=0, + ) + b_req_idx = torch.cat( + ( + model_input0.b_req_idx[0:req_num0], + model_input1.b_req_idx[0:req_num1], + ), + dim=0, + ) if req_num > 0: ( @@ -777,19 +780,55 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf model_input1, run_reqs1, ) = overlap_prepare_decode_inputs(req_objs=decode_reqs) - req_num0, req_num1 = len(run_reqs0), len(run_reqs1) - req_num = req_num0 + req_num1 - b_mtp_index_cpu0 = model_input0.b_mtp_index - b_mtp_index_cpu1 = model_input1.b_mtp_index + request_split = (len(decode_reqs) + 1) // 2 + real_request_num0 = request_split + real_request_num1 = len(decode_reqs) - request_split + req_num = real_request_num0 + real_request_num1 + spec_engine = self.decode_draft_engine with torch.cuda.stream(g_infer_context.get_overlap_stream()): + spec_plan = spec_engine.plan_decode( + model_input0=model_input0, + model_input1=model_input1, + decode_reqs=decode_reqs, + ) + ( + model_input0, + selected_row_mask_cpu0, + model_input1, + selected_row_mask_cpu1, + ) = spec_engine.prepare_decode_model_inputs( + model_input0=model_input0, + req_num0=real_request_num0, + model_input1=model_input1, + req_num1=real_request_num1, + plan=spec_plan, + ) model_output0, model_output1 = self.model.microbatch_overlap_decode(model_input0, model_input1) + + if selected_row_mask_cpu0 is not None: + selected_row_mask_cpu0.wait() + selected_rows0 = selected_row_mask_cpu0.tensor.tolist() + run_reqs0 = [req for req, selected in zip(run_reqs0, selected_rows0) if selected] + if selected_row_mask_cpu1 is not None: + selected_row_mask_cpu1.wait() + selected_rows1 = selected_row_mask_cpu1.tensor.tolist() + run_reqs1 = [req for req, selected in zip(run_reqs1, selected_rows1) if selected] + + verify_row_num0 = model_input0.batch_size + verify_row_num1 = model_input1.batch_size + verify_row_num = verify_row_num0 + verify_row_num1 logits0 = model_output0.logits logits1 = model_output1.logits run_reqs = run_reqs0 + run_reqs1 if req_num > 0: - logits = torch.empty((req_num, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) - logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) - logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1[0:req_num1, :], non_blocking=True) + assert len(run_reqs) == verify_row_num + logits = torch.empty( + (verify_row_num, logits0.shape[1]), + dtype=logits0.dtype, + device=logits0.device, + ) + logits[:verify_row_num0, :].copy_(logits0, non_blocking=True) + logits[verify_row_num0:, :].copy_(logits1, non_blocking=True) next_token_ids, next_token_logprobs = sample(logits, run_reqs, self.eos_id) next_token_ranks = self._get_next_token_ranks(logits, next_token_ids) ( @@ -798,31 +837,21 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf next_token_ranks_cpu, ) = self._async_copy_next_token_infos_to_pin_mem(next_token_ids, next_token_logprobs, next_token_ranks) - b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) - b_mtp_index_cpu = torch.cat((b_mtp_index_cpu0[0:req_num0], b_mtp_index_cpu1[0:req_num1]), dim=0) - b_req_mtp_start_loc = [ - index for index, mtp_index in enumerate(b_mtp_index_cpu.tolist()) if mtp_index == 0 - ] - b_req_mtp_start_loc = g_pin_mem_manager.gen_from_list( - key="b_req_mtp_start_loc", - data=b_req_mtp_start_loc, - dtype=torch.int32, - ).cuda(non_blocking=True) - - b_mtp_index = ( - torch.cat( - (model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), - dim=0, - ) - if self.is_linear_att_mixed_model - else None + b_req_idx = torch.cat((model_input0.b_req_idx, model_input1.b_req_idx), dim=0) + b_mtp_index = torch.cat( + (model_input0.b_mtp_index, model_input1.b_mtp_index), + dim=0, + ) + b_req_mtp_start_loc = gen_b_req_mtp_start_loc( + b_mtp_index=b_mtp_index, + num_reqs=req_num, ) mtp_accept_len, accepted_index = mtp_utils.verify_mtp_tokens( backend=self, next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, - b_mtp_index=b_mtp_index, + b_mtp_index=b_mtp_index if self.is_linear_att_mixed_model else None, ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", @@ -840,15 +869,19 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf verify_event = torch.cuda.Event() verify_event.record() + target_next_token_ids0 = next_token_ids[:verify_row_num0] + target_next_token_ids1 = next_token_ids[verify_row_num0:] proposal = self.decode_draft_engine.propose_next_overlap( target_model_input0=model_input0, target_model_output0=model_output0, + target_next_token_ids0=target_next_token_ids0, target_model_input1=model_input1, target_model_output1=model_output1, - target_next_token_ids=next_token_ids, - real_verify_rows0=req_num0, - real_verify_rows1=req_num1, + target_next_token_ids1=target_next_token_ids1, + real_request_num0=real_request_num0, + real_request_num1=real_request_num1, accept_len=mtp_accept_len, + draft_step=spec_plan.draft_step, ) if req_num > 0: mtp_utils.scatter_mtp_next_tokens( @@ -883,9 +916,13 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf event_pack.notify_forward_and_wait_post_handle() sync_event.synchronize() - mem_indexes_cpu = torch.cat( - (model_input0.mem_indexes_cpu[0:req_num0], model_input1.mem_indexes_cpu[0:req_num1]), dim=0 + spec_engine.update_planner_statics( + plan=spec_plan, + proposal=proposal, + req_num=req_num, + accept_lengths_cpu=mtp_accept_len_cpu, ) + mem_indexes_cpu = torch.cat((model_input0.mem_indexes_cpu, model_input1.mem_indexes_cpu), dim=0) proposal.extra_mem_indexes_cpu.append( MtpMemIndexesToFree( mem_indexes_cpu=mem_indexes_cpu, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index 8bae91971e..054e21a5ab 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -1,17 +1,22 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, List, Optional, Tuple import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner import ( - BaseDpOverlapPlanner, - build_dp_overlap_planner, +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers import ( + build_dp_overlap_spec_proposer, ) -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers import build_dp_overlap_spec_proposer -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( + BaseDpOverlapProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine +from lightllm.server.router.model_infer.mtp_speculative.planner import SpecDecodePlan +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + SpecProposal, +) +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor if TYPE_CHECKING: from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend @@ -20,13 +25,101 @@ class DPOverlapSpecEngine: """双 microbatch overlap draft 流程使用的 DP MTP engine。""" - def __init__(self, backend: ModeBackend, spec_mode: str, enable_dynmaic_mtp: bool) -> None: + def __init__( + self, + backend: ModeBackend, + spec_mode: str, + enable_dynmaic_mtp: bool, + common_engine: SpecEngine, + ) -> None: + self.common_engine = common_engine self.proposer: BaseDpOverlapProposer = build_dp_overlap_spec_proposer( spec_mode=spec_mode, backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp, ) - self.planner: BaseDpOverlapPlanner = build_dp_overlap_planner(backend=backend) + + def plan_decode( + self, + model_input0: ModelInput, + model_input1: ModelInput, + decode_reqs: List, + ) -> SpecDecodePlan: + """Use the common planner for the combined two-microbatch layout.""" + + return self.common_engine.planner.plan( + decode_reqs=decode_reqs, + origin_batch_size=model_input0.batch_size + model_input1.batch_size, + ) + + def prepare_decode_model_inputs( + self, + model_input0: ModelInput, + req_num0: int, + model_input1: ModelInput, + req_num1: int, + plan: SpecDecodePlan, + ) -> Tuple[ModelInput, Optional[AsyncPinnedCpuTensor], ModelInput, Optional[AsyncPinnedCpuTensor],]: + """Split the LightSpec verify budget and compact both microbatches. + + The combined budget is split in proportion to each side's real request + count, so every request receives a similar average number of verify + rows. Capacity overflow is transferred to the other side, and every + request still keeps at least its target row. + """ + + origin_batch_size = model_input0.batch_size + model_input1.batch_size + assert plan.origin_batch_size == origin_batch_size + if plan.dynamic_batch_size == origin_batch_size: + return model_input0, None, model_input1, None + + assert req_num0 == req_num1 or req_num0 == req_num1 + 1 + req_num = req_num0 + req_num1 + assert req_num > 0 + max_batch_size0 = min(model_input0.batch_size, req_num0 * (plan.pre_draft_step + 1)) + max_batch_size1 = min(model_input1.batch_size, req_num1 * (plan.pre_draft_step + 1)) + + # 先计算每个请求平均可分到的 verify 行数,再按请求数分配给第一侧。 + avg_verify_rows_per_req = plan.dynamic_batch_size / req_num + expected_batch_size0 = int(avg_verify_rows_per_req * req_num0 + 0.5) + min_batch_size0 = max(req_num0, plan.dynamic_batch_size - max_batch_size1) + max_batch_size0 = min(max_batch_size0, plan.dynamic_batch_size - req_num1) + dynamic_batch_size0 = min(max(expected_batch_size0, min_batch_size0), max_batch_size0) + dynamic_batch_size1 = plan.dynamic_batch_size - dynamic_batch_size0 + + assert req_num0 <= dynamic_batch_size0 <= max_batch_size0 + assert req_num1 <= dynamic_batch_size1 <= max_batch_size1 + + plan0 = SpecDecodePlan( + origin_batch_size=model_input0.batch_size, + dynamic_batch_size=dynamic_batch_size0, + draft_step=plan.draft_step, + pre_draft_step=plan.pre_draft_step, + all_reqs_have_proposals=plan.all_reqs_have_proposals, + ) + plan1 = SpecDecodePlan( + origin_batch_size=model_input1.batch_size, + dynamic_batch_size=dynamic_batch_size1, + draft_step=plan.draft_step, + pre_draft_step=plan.pre_draft_step, + all_reqs_have_proposals=plan.all_reqs_have_proposals, + ) + model_input0, selected_row_mask_cpu0 = self.common_engine.prepare_decode_model_input( + model_input=model_input0, + req_num=req_num0, + plan=plan0, + ) + model_input1, selected_row_mask_cpu1 = self.common_engine.prepare_decode_model_input( + model_input=model_input1, + req_num=req_num1, + plan=plan1, + ) + return ( + model_input0, + selected_row_mask_cpu0, + model_input1, + selected_row_mask_cpu1, + ) def fill_draft_model_kv_state_overlap( self, @@ -50,26 +143,44 @@ def propose_next_overlap( self, target_model_input0: ModelInput, # batch_size = verify_batch_size0 target_model_output0: ModelOutput, # logits: [verify_batch_size0, vocab_size] + target_next_token_ids0: torch.Tensor, # [verify_batch_size0] target_model_input1: ModelInput, # batch_size = verify_batch_size1 target_model_output1: ModelOutput, # logits: [verify_batch_size1, vocab_size] - target_next_token_ids: torch.Tensor, # [real_verify_rows0 + real_verify_rows1] - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, # [verify_batch_size1] + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, # [real_req_num0 + real_req_num1] + draft_step: int, ) -> SpecProposal: - real_verify_rows = real_verify_rows0 + real_verify_rows1 - assert target_next_token_ids.shape == (real_verify_rows,) + assert target_next_token_ids0.shape == (target_model_input0.batch_size,) + assert target_next_token_ids1.shape == (target_model_input1.batch_size,) + assert accept_len.shape == (real_request_num0 + real_request_num1,) return self.proposer.propose_next_overlap( target_model_input0=target_model_input0, target_model_output0=target_model_output0, + target_next_token_ids0=target_next_token_ids0, target_model_input1=target_model_input1, target_model_output1=target_model_output1, - target_next_token_ids=target_next_token_ids, - real_verify_rows0=real_verify_rows0, - real_verify_rows1=real_verify_rows1, + target_next_token_ids1=target_next_token_ids1, + real_request_num0=real_request_num0, + real_request_num1=real_request_num1, accept_len=accept_len, - draft_step=self.planner.get_draft_step(), + draft_step=draft_step, + ) + + def update_planner_statics( + self, + plan: SpecDecodePlan, + proposal: SpecProposal, + req_num: int, + accept_lengths_cpu: torch.Tensor, + ) -> None: + self.common_engine.update_planner_statics( + plan=plan, + proposal=proposal, + req_num=req_num, + accept_lengths_cpu=accept_lengths_cpu, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py deleted file mode 100644 index e85ee8cb34..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/__init__.py +++ /dev/null @@ -1,14 +0,0 @@ -from typing import TYPE_CHECKING - -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner.base import BaseDpOverlapPlanner -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner.fixed import FixedDpOverlapPlanner - -if TYPE_CHECKING: - from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend - - -def build_dp_overlap_planner(*, backend: "ModeBackend") -> BaseDpOverlapPlanner: - return FixedDpOverlapPlanner(draft_step=backend.max_draft_step) - - -__all__ = ["BaseDpOverlapPlanner", "FixedDpOverlapPlanner", "build_dp_overlap_planner"] diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/base.py deleted file mode 100644 index a809a28c6e..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/base.py +++ /dev/null @@ -1,11 +0,0 @@ -from abc import ABC, abstractmethod - - -class BaseDpOverlapPlanner(ABC): - """DP overlap draft 配置的基础规划接口。""" - - @abstractmethod - def get_draft_step(self) -> int: - """返回当前 overlap proposal 使用的 draft step。""" - - raise NotImplementedError diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/fixed.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/fixed.py deleted file mode 100644 index 3c96f9a0fc..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_planner/fixed.py +++ /dev/null @@ -1,11 +0,0 @@ -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner.base import BaseDpOverlapPlanner - - -class FixedDpOverlapPlanner(BaseDpOverlapPlanner): - """为 DP overlap 固定 shape decode 返回启动时配置的 draft step。""" - - def __init__(self, draft_step: int) -> None: - self.draft_step = int(draft_step) - - def get_draft_step(self) -> int: - return self.draft_step diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py index 2cc880b647..bdc5c757fe 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -4,7 +4,9 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + SpecProposal, +) if TYPE_CHECKING: from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend @@ -36,33 +38,15 @@ def propose_next_overlap( self, target_model_input0: ModelInput, # batch_size = padded_verify_batch_size0 target_model_output0: ModelOutput, # logits: [padded_verify_batch_size0, vocab_size] + target_next_token_ids0: torch.Tensor, # [padded_verify_batch_size0] target_model_input1: ModelInput, # batch_size = padded_verify_batch_size1 target_model_output1: ModelOutput, # logits: [padded_verify_batch_size1, vocab_size] - target_next_token_ids: torch.Tensor, # [real_verify_rows0 + real_verify_rows1] - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, # [padded_verify_batch_size1] + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, # [real_req_num0 + real_req_num1] draft_step: int, ) -> SpecProposal: """Generate one proposal from two DP-overlapped decode microbatches.""" raise NotImplementedError - - @staticmethod - def _build_padded_next_token_ids( - target_next_token_ids: torch.Tensor, - batch_size: int, - real_verify_rows: int, - device: torch.device, - source_start: int, - ) -> torch.Tensor: - """将合并后的真实 target token 拆分并补齐到一个 microbatch 的 verify shape。""" - - assert 0 <= real_verify_rows <= batch_size - padded_token_ids = torch.zeros((batch_size,), dtype=torch.int64, device=device) - if real_verify_rows > 0: - padded_token_ids[:real_verify_rows].copy_( - target_next_token_ids[source_start : source_start + real_verify_rows], - non_blocking=True, - ) - return padded_token_ids diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index e5a7e287c8..2f58a95794 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -1,12 +1,16 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( + BaseDpOverlapProposer, +) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, propose_next_dp_eagle_autoregressive_overlap, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + EagleSpecProposal, +) class DpOverlapEagle3Proposer(BaseDpOverlapProposer): @@ -38,11 +42,12 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids: torch.Tensor, - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: @@ -50,11 +55,12 @@ def propose_next_overlap( proposer=self, target_model_input0=target_model_input0, target_model_output0=target_model_output0, + target_next_token_ids0=target_next_token_ids0, target_model_input1=target_model_input1, target_model_output1=target_model_output1, - target_next_token_ids=target_next_token_ids, - real_verify_rows0=real_verify_rows0, - real_verify_rows1=real_verify_rows1, + target_next_token_ids1=target_next_token_ids1, + real_request_num0=real_request_num0, + real_request_num1=real_request_num1, accept_len=accept_len, draft_step=draft_step, map_draft_token_ids=self._map_draft_token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index c2250d5d51..546d25b765 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -1,12 +1,16 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( + BaseDpOverlapProposer, +) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, - propose_next_dp_eagle_fixed_layout_overlap, + propose_next_dp_no_att_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + EagleSpecProposal, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class DpOverlapEagleNoAttProposer(BaseDpOverlapProposer): @@ -35,24 +39,27 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids: torch.Tensor, - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: - return propose_next_dp_eagle_fixed_layout_overlap( + return propose_next_dp_no_att_overlap( proposer=self, target_model_input0=target_model_input0, target_model_output0=target_model_output0, + target_next_token_ids0=target_next_token_ids0, target_model_input1=target_model_input1, target_model_output1=target_model_output1, - target_next_token_ids=target_next_token_ids, - real_verify_rows0=real_verify_rows0, - real_verify_rows1=real_verify_rows1, + target_next_token_ids1=target_next_token_ids1, + real_request_num0=real_request_num0, + real_request_num1=real_request_num1, accept_len=accept_len, draft_step=draft_step, + get_draft_model=lambda _: self.backend.draft_models[0], map_draft_token_ids=lambda token_ids: token_ids, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py index de81db3e6f..1b818b5f3c 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py @@ -2,16 +2,29 @@ from __future__ import annotations +import copy from typing import Callable import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import ( + gen_mtp_new_input_ids, +) +from lightllm.common.basemodel.triton_kernel.mtp_utils import gen_b_req_mtp_start_loc +from lightllm.common.basemodel.triton_kernel.select_mtp_rows import ( + select_accepted_tail_rows, +) from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( + BaseDpOverlapProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + MtpMemIndexesToFree, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + EagleSpecProposal, +) from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @@ -37,12 +50,18 @@ def _prepare_eagle_prefill_inputs( model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens -def generate_eagle_token_ids( +def generate_eagle_tokens( proposer: BaseDpOverlapProposer, model_output: ModelOutput, map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> torch.Tensor: - return map_draft_token_ids(proposer.backend._gen_argmax_token_ids(model_output)) +) -> tuple[torch.Tensor, torch.Tensor | None]: + if proposer.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = proposer.backend._gen_argmax_token_ids_and_prob(model_output) + return map_draft_token_ids(draft_token_ids), draft_token_probs.float() + return ( + map_draft_token_ids(proposer.backend._gen_argmax_token_ids(model_output)), + None, + ) def prepare_eagle_verify_decode_input( @@ -81,136 +100,95 @@ def fill_dp_eagle_draft_model_kv_state_overlap( proposer.backend.draft_models[0]._microbatch_overlap_prefill_cuda(target_model_input0, target_model_input1) -def pad_dp_step_mem_indexes( - real_mem_indexes: torch.Tensor, - request_capacity: int, - hold_mem_index: int, -) -> torch.Tensor: - padded = torch.full( - (request_capacity,), - hold_mem_index, - dtype=real_mem_indexes.dtype, - device=real_mem_indexes.device, - ) - padded[: real_mem_indexes.shape[0]].copy_(real_mem_indexes) - return padded +def get_dp_overlap_req_start_rows(b_mtp_index: torch.Tensor, request_num: int) -> torch.Tensor: + """根据动态 verify 布局恢复请求起始行。""" + + request_num = int(request_num) + if request_num == 0: + assert b_mtp_index.numel() == 0 + return torch.empty((0,), dtype=torch.int32, device=b_mtp_index.device) + if b_mtp_index.is_cuda: + return gen_b_req_mtp_start_loc( + b_mtp_index=b_mtp_index, + num_reqs=request_num, + ) + req_start_rows = torch.nonzero(b_mtp_index == 0, as_tuple=False).flatten().to(dtype=torch.int32) + assert req_start_rows.shape == (request_num,) + return req_start_rows def prepare_dp_eagle_overlap_decode_inputs( - proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_input1: ModelInput, - target_next_token_ids: torch.Tensor, - real_verify_rows0: int, - real_verify_rows1: int, + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - """拆分真实 verify 结果,并补齐为两个 EAGLE overlap microbatch 的固定布局。""" - - verify_width = proposer.backend.max_draft_step + 1 - real_request_num0 = real_verify_rows0 // verify_width - real_request_num1 = real_verify_rows1 // verify_width - request_capacity0 = target_model_input0.batch_size // verify_width - request_capacity1 = target_model_input1.batch_size // verify_width + """恢复两个动态 verify microbatch 的请求起始行。""" assert accept_len.shape == (real_request_num0 + real_request_num1,) - - target_next_token_ids0 = proposer._build_padded_next_token_ids( - target_next_token_ids=target_next_token_ids, - batch_size=target_model_input0.batch_size, - real_verify_rows=real_verify_rows0, - device=target_model_input0.b_req_idx.device, - source_start=0, + accept_len0 = accept_len[:real_request_num0] + accept_len1 = accept_len[real_request_num0:] + req_start_rows0 = get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input0.b_mtp_index, + request_num=real_request_num0, ) - target_next_token_ids1 = proposer._build_padded_next_token_ids( - target_next_token_ids=target_next_token_ids, - batch_size=target_model_input1.batch_size, - real_verify_rows=real_verify_rows1, - device=target_model_input1.b_req_idx.device, - source_start=real_verify_rows0, - ) - - padded_accept_len0 = torch.ones( - (request_capacity0,), - dtype=torch.int32, - device=target_model_input0.b_req_idx.device, + req_start_rows1 = get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input1.b_mtp_index, + request_num=real_request_num1, ) - padded_accept_len1 = torch.ones( - (request_capacity1,), - dtype=torch.int32, - device=target_model_input1.b_req_idx.device, + return ( + accept_len0, + accept_len1, + req_start_rows0, + req_start_rows1, ) - if real_request_num0 > 0: - padded_accept_len0[:real_request_num0].copy_(accept_len[:real_request_num0]) - if real_request_num1 > 0: - padded_accept_len1[:real_request_num1].copy_( - accept_len[real_request_num0 : real_request_num0 + real_request_num1] - ) - - return target_next_token_ids0, target_next_token_ids1, padded_accept_len0, padded_accept_len1 def propose_next_dp_eagle_autoregressive_overlap( proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids: torch.Tensor, - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, draft_step: int, map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], ) -> EagleSpecProposal: """运行 DP EAGLE extend 后接单 token overlap decode 的 proposal 流程。""" - ( - target_next_token_ids0, - target_next_token_ids1, - accept_len0, - accept_len1, - ) = prepare_dp_eagle_overlap_decode_inputs( - proposer=proposer, + assert target_next_token_ids0.shape == (target_model_input0.batch_size,) + assert target_next_token_ids1.shape == (target_model_input1.batch_size,) + (accept_len0, accept_len1, req_start_rows0, req_start_rows1,) = prepare_dp_eagle_overlap_decode_inputs( target_model_input0=target_model_input0, target_model_input1=target_model_input1, - target_next_token_ids=target_next_token_ids, - real_verify_rows0=real_verify_rows0, - real_verify_rows1=real_verify_rows1, + real_request_num0=real_request_num0, + real_request_num1=real_request_num1, accept_len=accept_len, ) - verify_width = proposer.backend.max_draft_step + 1 - model_inputs = (target_model_input0, target_model_input1) + model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] model_outputs = (target_model_output0, target_model_output1) target_next_token_ids_by_batch = (target_next_token_ids0, target_next_token_ids1) - real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) accept_lens_by_batch = (accept_len0, accept_len1) + req_start_rows_by_batch = (req_start_rows0, req_start_rows1) position_deltas_by_batch = tuple(model_input.b_position_delta for model_input in model_inputs) max_kv_seq_lens_by_batch = tuple(model_input.max_kv_seq_len for model_input in model_inputs) - request_capacities_by_batch = [] - real_request_counts = [] + real_request_counts = (int(real_request_num0), int(real_request_num1)) accepted_tail_rows_by_batch = [] - for model_input, model_output, token_ids, real_verify_row_count, accept_len in zip( + for model_input, model_output, token_ids, req_start_rows, accept_len in zip( model_inputs, model_outputs, target_next_token_ids_by_batch, - real_verify_row_counts, + req_start_rows_by_batch, accept_lens_by_batch, ): - request_capacity = model_input.batch_size // verify_width - real_request_count = real_verify_row_count // verify_width - starts = torch.arange( - 0, - model_input.batch_size, - verify_width, - dtype=torch.int32, - device=token_ids.device, - ) - accepted_tail_rows = (starts + accept_len - 1).long() - request_capacities_by_batch.append(request_capacity) - real_request_counts.append(real_request_count) + accepted_tail_rows = (req_start_rows + accept_len - 1).long() accepted_tail_rows_by_batch.append(accepted_tail_rows) prepare_eagle_verify_decode_input( model_input=model_input, @@ -220,6 +198,15 @@ def propose_next_dp_eagle_autoregressive_overlap( total_real_request_count = sum(real_request_counts) proposal_token_ids = target_next_token_ids0.new_empty((total_real_request_count, draft_step)) + proposal_schedule_scores = ( + torch.empty( + (total_real_request_count, draft_step), + dtype=torch.float32, + device=target_next_token_ids0.device, + ) + if proposer.enable_dynmaic_mtp + else None + ) draft_model = proposer.backend.draft_models[0] extend_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) @@ -231,11 +218,20 @@ def propose_next_dp_eagle_autoregressive_overlap( draft_shared_seq_lens_by_batch = [] draft_shared_radix_node_ids_by_batch = [] proposal_row_offsets = (0, real_request_counts[0]) - for batch_index, (model_input, extend_output, accepted_tail_rows, real_request_count) in enumerate( - zip(model_inputs, extend_outputs, accepted_tail_rows_by_batch, real_request_counts) + for batch_index, (model_input, extend_output, accepted_tail_rows, real_request_count,) in enumerate( + zip( + model_inputs, + extend_outputs, + accepted_tail_rows_by_batch, + real_request_counts, + ) ): accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) - draft_token_ids = generate_eagle_token_ids(proposer, accepted_tail_output, map_draft_token_ids) + draft_token_ids, draft_token_probs = generate_eagle_tokens( + proposer, + accepted_tail_output, + map_draft_token_ids, + ) draft_token_ids_by_batch.append(draft_token_ids) draft_hiddens_by_batch.append(extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows)) draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) @@ -248,17 +244,21 @@ def propose_next_dp_eagle_autoregressive_overlap( proposal_token_ids[proposal_row_start : proposal_row_start + real_request_count, 0] = draft_token_ids[ :real_request_count ] + if proposal_schedule_scores is not None: + proposal_schedule_scores[ + proposal_row_start : proposal_row_start + real_request_count, 0 + ] = draft_token_probs[:real_request_count] if draft_step == 1: return EagleSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=[], - schedule_scores=None, + schedule_scores=proposal_schedule_scores, ) for batch_index, model_input in enumerate(model_inputs): model_input.is_prefill = False - model_input.batch_size = request_capacities_by_batch[batch_index] + model_input.batch_size = real_request_counts[batch_index] model_input.b_req_idx = draft_req_indices_by_batch[batch_index] model_input.b_mtp_index = torch.zeros_like(model_input.b_req_idx) model_input.b_seq_len = draft_seq_lens_by_batch[batch_index] @@ -274,7 +274,6 @@ def propose_next_dp_eagle_autoregressive_overlap( extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * (draft_step - 1)) extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids0.device, non_blocking=True) - hold_mem_index = proposer.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX for step in range(1, draft_step): mem_start = (step - 1) * total_real_request_count @@ -286,17 +285,17 @@ def propose_next_dp_eagle_autoregressive_overlap( real_mem_start += real_request_count model_input.input_ids = draft_token_ids_by_batch[batch_index] model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - model_input.mem_indexes = pad_dp_step_mem_indexes( - real_mem_indexes=real_mem_indexes, - request_capacity=request_capacities_by_batch[batch_index], - hold_mem_index=hold_mem_index, - ) + model_input.mem_indexes = real_mem_indexes model_input.max_kv_seq_len = max_kv_seq_lens_by_batch[batch_index] + step model_input.total_token_num = model_input.batch_size * model_input.max_kv_seq_len draft_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) for batch_index, draft_output in enumerate(draft_outputs): - draft_token_ids = generate_eagle_token_ids(proposer, draft_output, map_draft_token_ids) + draft_token_ids, draft_token_probs = generate_eagle_tokens( + proposer, + draft_output, + map_draft_token_ids, + ) draft_token_ids_by_batch[batch_index] = draft_token_ids draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden draft_seq_lens_by_batch[batch_index].add_(1) @@ -305,124 +304,136 @@ def propose_next_dp_eagle_autoregressive_overlap( proposal_token_ids[proposal_row_start : proposal_row_start + real_request_count, step] = draft_token_ids[ :real_request_count ] + if proposal_schedule_scores is not None: + proposal_schedule_scores[ + proposal_row_start : proposal_row_start + real_request_count, + step, + ] = draft_token_probs[:real_request_count] return EagleSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], - schedule_scores=None, + schedule_scores=proposal_schedule_scores, ) -def propose_next_dp_eagle_fixed_layout_overlap( +def propose_next_dp_no_att_overlap( proposer: BaseDpOverlapProposer, target_model_input0: ModelInput, target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids: torch.Tensor, - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, draft_step: int, + get_draft_model: Callable[[int], object], map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], ) -> EagleSpecProposal: - """以 expanded verify-row layout 运行 decode,返回按真实请求压缩的 proposal。""" + """将动态 verify 布局压缩为每请求一行,再执行双 microbatch No-Att draft。""" + + real_request_counts = (int(real_request_num0), int(real_request_num1)) + total_request_num = sum(real_request_counts) + proposal_token_ids = target_next_token_ids0.new_empty((total_request_num, draft_step)) + proposal_schedule_scores = ( + torch.empty( + (total_request_num, draft_step), + dtype=torch.float32, + device=target_next_token_ids0.device, + ) + if proposer.enable_dynmaic_mtp + else None + ) + if draft_step == 0: + return EagleSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=proposal_schedule_scores, + ) - ( - target_next_token_ids0, - target_next_token_ids1, - accept_len0, - accept_len1, - ) = prepare_dp_eagle_overlap_decode_inputs( - proposer=proposer, + assert target_next_token_ids0.shape == (target_model_input0.batch_size,) + assert target_next_token_ids1.shape == (target_model_input1.batch_size,) + (accept_len0, accept_len1, req_start_rows0, req_start_rows1,) = prepare_dp_eagle_overlap_decode_inputs( target_model_input0=target_model_input0, target_model_input1=target_model_input1, - target_next_token_ids=target_next_token_ids, - real_verify_rows0=real_verify_rows0, - real_verify_rows1=real_verify_rows1, + real_request_num0=real_request_num0, + real_request_num1=real_request_num1, accept_len=accept_len, ) - verify_width = proposer.backend.max_draft_step + 1 - model_inputs = (target_model_input0, target_model_input1) - real_verify_row_counts = (int(real_verify_rows0), int(real_verify_rows1)) - real_request_counts = tuple(row_count // verify_width for row_count in real_verify_row_counts) - request_capacities_by_batch = tuple(model_input.batch_size // verify_width for model_input in model_inputs) - total_real_request_count = sum(real_request_counts) + model_inputs = [target_model_input0, target_model_input1] + model_outputs = [target_model_output0, target_model_output1] + target_token_ids_by_batch = [target_next_token_ids0, target_next_token_ids1] + req_start_rows_by_batch = [req_start_rows0, req_start_rows1] + accept_lens_by_batch = [accept_len0, accept_len1] + draft_inputs = [] + draft_token_ids_by_batch = [] + draft_hiddens_by_batch = [] + for (model_input, model_output, token_ids, req_start_rows, batch_accept_len, request_num,) in zip( + model_inputs, + model_outputs, + target_token_ids_by_batch, + req_start_rows_by_batch, + accept_lens_by_batch, + real_request_counts, + ): + selected_rows = select_accepted_tail_rows( + b_req_mtp_start_loc=req_start_rows, + accept_len=batch_accept_len, + input_ids=token_ids, + hidden=model_output.mtp_collector.spec_hidden, + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, + b_seq_len=model_input.b_seq_len, + mem_indexes=model_input.mem_indexes, + b_shared_seq_len=model_input.b_shared_seq_len, + b_shared_radix_node_id=model_input.b_shared_radix_node_id, + b_position_delta=model_input.b_position_delta, + ) + draft_input = copy.copy(model_input) + draft_input.batch_size = request_num + draft_input.input_ids = selected_rows.input_ids + draft_input.b_req_idx = selected_rows.b_req_idx + draft_input.b_mtp_index = selected_rows.b_mtp_index + draft_input.b_seq_len = selected_rows.b_seq_len + draft_input.mem_indexes = selected_rows.mem_indexes + draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len + draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id + draft_input.b_position_delta = selected_rows.b_position_delta + draft_input.mem_indexes_cpu = None + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(request_num)] + draft_inputs.append(draft_input) + draft_token_ids_by_batch.append(selected_rows.input_ids) + draft_hiddens_by_batch.append(selected_rows.hidden) - proposal_token_ids = target_next_token_ids0.new_empty((total_real_request_count, draft_step)) proposal_row_offsets = (0, real_request_counts[0]) - accepted_tail_rows_by_batch = ( - torch.arange(0, target_model_input0.batch_size, verify_width, device=target_next_token_ids0.device) - + accept_len0 - - 1, - torch.arange(0, target_model_input1.batch_size, verify_width, device=target_next_token_ids1.device) - + accept_len1 - - 1, - ) - - extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * draft_step) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids0.device, non_blocking=True) - split = real_request_counts[0] * draft_step - extra_mem_indexes_by_batch = ( - extra_mem_indexes[:split], - extra_mem_indexes[split:], - ) - - draft_token_ids_by_batch = [target_next_token_ids0, target_next_token_ids1] - draft_hiddens_by_batch = [ - target_model_output0.mtp_collector.spec_hidden, - target_model_output1.mtp_collector.spec_hidden, - ] - draft_model = proposer.backend.draft_models[0] - hold_mem_index = proposer.backend.model.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX - for step in range(draft_step): - for batch_index, model_input in enumerate(model_inputs): - model_input.input_ids = draft_token_ids_by_batch[batch_index] - model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - - draft_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) + for batch_index, draft_input in enumerate(draft_inputs): + draft_input.input_ids = draft_token_ids_by_batch[batch_index] + draft_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - for batch_index, (model_input, draft_output) in enumerate(zip(model_inputs, draft_outputs)): - model_input.b_seq_len += 1 - model_input.max_kv_seq_len += 1 - - real_request_count = real_request_counts[batch_index] - mem_start = step * real_request_count - step_mem_indexes = extra_mem_indexes_by_batch[batch_index][mem_start : mem_start + real_request_count] - step_mem_indexes = pad_dp_step_mem_indexes( - real_mem_indexes=step_mem_indexes, - request_capacity=request_capacities_by_batch[batch_index], - hold_mem_index=hold_mem_index, - ) - model_input.mem_indexes = torch.cat( - [ - model_input.mem_indexes.view(-1, verify_width)[:, 1:], - step_mem_indexes.view(-1, 1), - ], - dim=1, - ).view(-1) - - draft_token_ids_by_batch[batch_index] = generate_eagle_token_ids( - proposer, - draft_output, - map_draft_token_ids, + draft_outputs = get_draft_model(step)._microbatch_overlap_decode_cuda(*draft_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + draft_token_ids, draft_token_probs = generate_eagle_tokens( + proposer=proposer, + model_output=draft_output, + map_draft_token_ids=map_draft_token_ids, ) + draft_token_ids_by_batch[batch_index] = draft_token_ids draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden - - for batch_index, draft_token_ids in enumerate(draft_token_ids_by_batch): - real_request_count = real_request_counts[batch_index] + request_num = real_request_counts[batch_index] proposal_row_start = proposal_row_offsets[batch_index] - proposal_token_ids[ - proposal_row_start : proposal_row_start + real_request_count, step - ] = draft_token_ids.index_select( - 0, - accepted_tail_rows_by_batch[batch_index][:real_request_count].long(), - ) + proposal_token_ids[proposal_row_start : proposal_row_start + request_num, step] = draft_token_ids + if proposal_schedule_scores is not None: + proposal_schedule_scores[ + proposal_row_start : proposal_row_start + request_num, + step, + ] = draft_token_probs return EagleSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], - schedule_scores=None, + extra_mem_indexes_cpu=[], + schedule_scores=proposal_schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index dc0e1609f3..3a5c45f981 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -1,12 +1,16 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( + BaseDpOverlapProposer, +) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( fill_dp_eagle_draft_model_kv_state_overlap, - propose_next_dp_eagle_fixed_layout_overlap, + propose_next_dp_eagle_autoregressive_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + EagleSpecProposal, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal class DpOverlapEagleWithAttProposer(BaseDpOverlapProposer): @@ -35,23 +39,25 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids: torch.Tensor, - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: - return propose_next_dp_eagle_fixed_layout_overlap( + return propose_next_dp_eagle_autoregressive_overlap( proposer=self, target_model_input0=target_model_input0, target_model_output0=target_model_output0, + target_next_token_ids0=target_next_token_ids0, target_model_input1=target_model_input1, target_model_output1=target_model_output1, - target_next_token_ids=target_next_token_ids, - real_verify_rows0=real_verify_rows0, - real_verify_rows1=real_verify_rows1, + target_next_token_ids1=target_next_token_ids1, + real_request_num0=real_request_num0, + real_request_num1=real_request_num1, accept_len=accept_len, draft_step=draft_step, map_draft_token_ids=lambda token_ids: token_ids, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 81ca39de3d..5c0607c082 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -1,10 +1,15 @@ -import copy - import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( + BaseDpOverlapProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( + propose_next_dp_no_att_overlap, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + VanillaSpecProposal, +) class DpOverlapVanillaNoAttProposer(BaseDpOverlapProposer): @@ -25,65 +30,32 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids: torch.Tensor, - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: - verify_width = self.backend.max_draft_step + 1 - real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) - req_num_by_batch = tuple(row_count // verify_width for row_count in real_verify_rows) - assert accept_len.shape == (sum(req_num_by_batch),) - - target_next_token_ids0 = self._build_padded_next_token_ids( - target_next_token_ids=target_next_token_ids, - batch_size=target_model_input0.batch_size, - real_verify_rows=real_verify_rows0, - device=target_model_input0.b_req_idx.device, - source_start=0, - ) - target_next_token_ids1 = self._build_padded_next_token_ids( - target_next_token_ids=target_next_token_ids, - batch_size=target_model_input1.batch_size, - real_verify_rows=real_verify_rows1, - device=target_model_input1.b_req_idx.device, - source_start=real_verify_rows0, - ) - req_start_rows = ( - torch.arange(0, real_verify_rows0, verify_width, device=target_next_token_ids0.device), - torch.arange(0, real_verify_rows1, verify_width, device=target_next_token_ids1.device), + proposal = propose_next_dp_no_att_overlap( + proposer=self, + target_model_input0=target_model_input0, + target_model_output0=target_model_output0, + target_next_token_ids0=target_next_token_ids0, + target_model_input1=target_model_input1, + target_model_output1=target_model_output1, + target_next_token_ids1=target_next_token_ids1, + real_request_num0=real_request_num0, + real_request_num1=real_request_num1, + accept_len=accept_len, + draft_step=draft_step, + get_draft_model=lambda step: self.backend.draft_models[step], + map_draft_token_ids=lambda token_ids: token_ids, ) - accepted_tail_rows = ( - req_start_rows[0] + accept_len[: req_num_by_batch[0]] - 1, - req_start_rows[1] + accept_len[req_num_by_batch[0] : sum(req_num_by_batch)] - 1, - ) - model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] - draft_token_ids = [target_next_token_ids0, target_next_token_ids1] - draft_hiddens = [ - target_model_output0.mtp_collector.spec_hidden, - target_model_output1.mtp_collector.spec_hidden, - ] - proposal_token_ids = target_next_token_ids0.new_empty((sum(req_num_by_batch), draft_step)) - req_offset = req_num_by_batch[0] - - for step in range(draft_step): - for batch_index, model_input in enumerate(model_inputs): - model_input.input_ids = draft_token_ids[batch_index] - model_input.mtp_draft_input_hiddens = draft_hiddens[batch_index] - - draft_outputs = self.backend.draft_models[step]._microbatch_overlap_decode_cuda(*model_inputs) - for batch_index, draft_output in enumerate(draft_outputs): - draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden - draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) - - proposal_token_ids[:req_offset, step] = draft_token_ids[0].index_select(0, accepted_tail_rows[0].long()) - proposal_token_ids[req_offset:, step] = draft_token_ids[1].index_select(0, accepted_tail_rows[1].long()) - return VanillaSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=None, + token_ids=proposal.token_ids, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, + schedule_scores=proposal.schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 895c55dbdf..1a6b8dec4d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -3,10 +3,21 @@ import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids -from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import ( + gen_mtp_new_input_ids, +) +from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import ( + overlay_chained_mtp_decode_input, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( + BaseDpOverlapProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( + get_dp_overlap_req_start_rows, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + VanillaSpecProposal, +) from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @@ -52,39 +63,32 @@ def propose_next_overlap( self, target_model_input0: ModelInput, target_model_output0: ModelOutput, + target_next_token_ids0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, - target_next_token_ids: torch.Tensor, - real_verify_rows0: int, - real_verify_rows1: int, + target_next_token_ids1: torch.Tensor, + real_request_num0: int, + real_request_num1: int, accept_len: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: assert draft_step == self.backend.max_draft_step assert len(self.backend.draft_models) == draft_step - verify_width = self.backend.max_draft_step + 1 - real_verify_rows = (int(real_verify_rows0), int(real_verify_rows1)) - req_num_by_batch = tuple(row_count // verify_width for row_count in real_verify_rows) + req_num_by_batch = (int(real_request_num0), int(real_request_num1)) assert accept_len.shape == (sum(req_num_by_batch),) - target_next_token_ids0 = self._build_padded_next_token_ids( - target_next_token_ids=target_next_token_ids, - batch_size=target_model_input0.batch_size, - real_verify_rows=real_verify_rows0, - device=target_model_input0.b_req_idx.device, - source_start=0, - ) - target_next_token_ids1 = self._build_padded_next_token_ids( - target_next_token_ids=target_next_token_ids, - batch_size=target_model_input1.batch_size, - real_verify_rows=real_verify_rows1, - device=target_model_input1.b_req_idx.device, - source_start=real_verify_rows0, - ) + assert target_next_token_ids0.shape == (target_model_input0.batch_size,) + assert target_next_token_ids1.shape == (target_model_input1.batch_size,) req_start_rows = ( - torch.arange(0, real_verify_rows0, verify_width, device=target_next_token_ids0.device), - torch.arange(0, real_verify_rows1, verify_width, device=target_next_token_ids1.device), + get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input0.b_mtp_index, + request_num=real_request_num0, + ), + get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input1.b_mtp_index, + request_num=real_request_num1, + ), ) accept_len_by_batch = ( accept_len[: req_num_by_batch[0]], @@ -98,6 +102,15 @@ def propose_next_overlap( target_model_output1.mtp_collector.spec_hidden, ] proposal_token_ids = target_next_token_ids0.new_empty((sum(req_num_by_batch), draft_step)) + proposal_schedule_scores = ( + torch.empty( + (sum(req_num_by_batch), draft_step), + dtype=torch.float32, + device=target_next_token_ids0.device, + ) + if self.enable_dynmaic_mtp + else None + ) req_offset = req_num_by_batch[0] for step in range(draft_step): @@ -108,7 +121,21 @@ def propose_next_overlap( draft_outputs = self.backend.draft_models[step]._microbatch_overlap_decode_cuda(*model_inputs) for batch_index, draft_output in enumerate(draft_outputs): draft_hiddens[batch_index] = draft_output.mtp_collector.spec_hidden - draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) + if self.enable_dynmaic_mtp: + draft_token_ids[batch_index], draft_token_probs = self.backend._gen_argmax_token_ids_and_prob( + draft_output + ) + selected_token_probs = draft_token_probs.index_select( + 0, + accepted_tail_rows[batch_index].long(), + ) + proposal_row_start = 0 if batch_index == 0 else req_offset + proposal_schedule_scores[ + proposal_row_start : proposal_row_start + req_num_by_batch[batch_index], + step, + ] = selected_token_probs + else: + draft_token_ids[batch_index] = self.backend._gen_argmax_token_ids(draft_output) proposal_token_ids[:req_offset, step] = draft_token_ids[0].index_select(0, accepted_tail_rows[0].long()) proposal_token_ids[req_offset:, step] = draft_token_ids[1].index_select(0, accepted_tail_rows[1].long()) @@ -125,7 +152,7 @@ def propose_next_overlap( return VanillaSpecProposal( token_ids=proposal_token_ids, extra_mem_indexes_cpu=[], - schedule_scores=None, + schedule_scores=proposal_schedule_scores, ) def _prepare_mtp_prefill_inputs( diff --git a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py index 2cd781ce49..b25c62b9b7 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -68,11 +68,13 @@ def __init__( # The current verify width is bounded by the proposal built last time. self.pre_draft_step = self.max_draft_step - # 只有非 overlap DP 下的变长 LightSpec 需要跨 rank 对齐 draft 深度。 + # DP 下的变长 LightSpec 需要跨 rank 对齐 draft 深度。 + # overlap engine 与普通 engine 共享这一 planner,因此全局只会 + # 创建一组通信资源。 self._draft_step_group = None self._draft_step_tensor = None self._draft_step_stream = None - if backend.args.dp > 1 and not backend.enable_decode_microbatch_overlap and len(self.draft_steps) > 1: + if backend.args.dp > 1 and len(self.draft_steps) > 1: self._draft_step_group = dist.new_group( ranks=list(range(get_global_world_size())), backend="nccl", diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 03bf193f68..2bc7e809e3 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -4,19 +4,31 @@ import pytest import torch -from lightllm.server.router.model_infer.mode_backend.dp_backend import impl as dp_backend_impl -from lightllm.server.router.model_infer.mode_backend.dp_backend.impl import DPChunkedPrefillBackend +from lightllm.server.router.model_infer.mode_backend.dp_backend import ( + impl as dp_backend_impl, +) +from lightllm.server.router.model_infer.mode_backend.dp_backend.impl import ( + DPChunkedPrefillBackend, +) from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend -from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ChunkedPrefillBackend -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine +from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ( + ChunkedPrefillBackend, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import ( + DPOverlapSpecEngine, +) from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine -from lightllm.server.router.model_infer.mtp_speculative.planner import lightspec as lightspec_planner_impl +from lightllm.server.router.model_infer.mtp_speculative.planner import ( + lightspec as lightspec_planner_impl, +) from lightllm.server.router.model_infer.mtp_speculative.planner import ( FixedSpecPlanner, LightSpecPlanner, SpecDecodePlan, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + SpecProposal, +) def test_dp_backend_reuses_common_engine_outside_overlap(): @@ -39,6 +51,7 @@ def test_dp_backend_reuses_common_engine_outside_overlap(): assert not dp_backend.spec_engine.proposer.enable_dynmaic_mtp assert type(dp_backend.spec_engine.planner) is FixedSpecPlanner assert type(dp_backend.dp_overlap_spec_engine) is DPOverlapSpecEngine + assert dp_backend.dp_overlap_spec_engine.common_engine is dp_backend.spec_engine assert dp_backend.prefill_draft_engine is dp_backend.spec_engine assert dp_backend.decode_draft_engine is dp_backend.dp_overlap_spec_engine assert not issubclass(DPOverlapSpecEngine, SpecEngine) @@ -148,12 +161,21 @@ def test_dp_backend_does_not_build_group_for_fixed_planner(monkeypatch): assert not hasattr(backend.spec_engine.planner, "_draft_step_group") -def test_dp_backend_does_not_build_group_for_overlap_decode(monkeypatch): +def test_dp_backend_overlap_decode_reuses_common_lightspec_group(monkeypatch): + created_group = object() + created_stream = object() + real_torch_zeros = torch.zeros + + monkeypatch.setattr(lightspec_planner_impl, "get_global_world_size", lambda: 2) + monkeypatch.setattr(lightspec_planner_impl.dist, "new_group", lambda **kwargs: created_group) monkeypatch.setattr( - lightspec_planner_impl.dist, - "new_group", - lambda **kwargs: pytest.fail(f"unexpected NCCL group: {kwargs}"), + lightspec_planner_impl.torch, + "zeros", + lambda *args, **kwargs: real_torch_zeros( + *args, **{key: value for key, value in kwargs.items() if key != "device"} + ), ) + monkeypatch.setattr(lightspec_planner_impl.torch.cuda, "Stream", lambda: created_stream) backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) backend.args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=True, dp=2) @@ -166,8 +188,10 @@ def test_dp_backend_does_not_build_group_for_overlap_decode(monkeypatch): backend.init_spec_engine() - assert backend.spec_engine.planner._draft_step_group is None + assert backend.spec_engine.planner._draft_step_group is created_group + assert backend.spec_engine.planner._draft_step_stream is created_stream assert backend.decode_draft_engine is backend.dp_overlap_spec_engine + assert backend.decode_draft_engine.common_engine.planner is backend.spec_engine.planner def test_dp_backend_keeps_dynamic_planning_before_global_reduction(): @@ -272,7 +296,15 @@ def propose_next(self, **kwargs): backend.decode_mtp(event_pack=event_pack, decode_reqs=[]) - assert calls == ["plan", "prepare", "propose", "post_wait", "forward_wait", "free", "pre_post"] + assert calls == [ + "plan", + "prepare", + "propose", + "post_wait", + "forward_wait", + "free", + "pre_post", + ] def test_dp_overlap_engine_delegates_raw_verify_layout_to_proposer(): @@ -281,41 +313,121 @@ def test_dp_overlap_engine_delegates_raw_verify_layout_to_proposer(): class _Proposer: def propose_next_overlap(self, **kwargs): calls.update(kwargs) - return SpecProposal(token_ids=kwargs["target_next_token_ids"].new_empty((3, 7))) + return SpecProposal(token_ids=kwargs["target_next_token_ids0"].new_empty((3, 7))) engine = DPOverlapSpecEngine.__new__(DPOverlapSpecEngine) engine.proposer = _Proposer() - engine.planner = SimpleNamespace(get_draft_step=lambda: 7) - model_input0 = SimpleNamespace(batch_size=16) + model_input0 = SimpleNamespace(batch_size=8) model_input1 = SimpleNamespace(batch_size=16) model_output0 = SimpleNamespace() model_output1 = SimpleNamespace() - target_next_token_ids = torch.arange(24, dtype=torch.int64) + target_next_token_ids0 = torch.arange(8, dtype=torch.int64) + target_next_token_ids1 = torch.arange(8, 24, dtype=torch.int64) accept_len = torch.tensor([2, 3, 4], dtype=torch.int32) proposal = engine.propose_next_overlap( target_model_input0=model_input0, target_model_output0=model_output0, + target_next_token_ids0=target_next_token_ids0, target_model_input1=model_input1, target_model_output1=model_output1, - target_next_token_ids=target_next_token_ids, - real_verify_rows0=8, - real_verify_rows1=16, + target_next_token_ids1=target_next_token_ids1, + real_request_num0=1, + real_request_num1=2, accept_len=accept_len, + draft_step=7, ) assert calls["target_model_input0"] is model_input0 assert calls["target_model_input1"] is model_input1 assert calls["target_model_output0"] is model_output0 assert calls["target_model_output1"] is model_output1 - assert calls["target_next_token_ids"] is target_next_token_ids - assert calls["real_verify_rows0"] == 8 - assert calls["real_verify_rows1"] == 16 + assert calls["target_next_token_ids0"] is target_next_token_ids0 + assert calls["target_next_token_ids1"] is target_next_token_ids1 + assert calls["real_request_num0"] == 1 + assert calls["real_request_num1"] == 2 assert calls["accept_len"] is accept_len assert calls["draft_step"] == 7 assert proposal.token_ids.shape == (3, 7) +def test_dp_overlap_engine_distributes_dynamic_verify_budget_by_request_count(): + prepare_calls = [] + + class _CommonEngine: + def prepare_decode_model_input(self, **kwargs): + prepare_calls.append(kwargs) + model_input = kwargs["model_input"] + model_input.batch_size = kwargs["plan"].dynamic_batch_size + return model_input, f"mask{len(prepare_calls)}" + + engine = DPOverlapSpecEngine.__new__(DPOverlapSpecEngine) + engine.common_engine = _CommonEngine() + model_input0 = SimpleNamespace(batch_size=12) + model_input1 = SimpleNamespace(batch_size=8) + plan = SpecDecodePlan( + origin_batch_size=20, + dynamic_batch_size=11, + draft_step=2, + pre_draft_step=3, + ) + + compacted_input0, mask0, compacted_input1, mask1 = engine.prepare_decode_model_inputs( + model_input0=model_input0, + req_num0=3, + model_input1=model_input1, + req_num1=2, + plan=plan, + ) + + assert compacted_input0.batch_size == 7 + assert compacted_input1.batch_size == 4 + assert mask0 == "mask1" + assert mask1 == "mask2" + assert prepare_calls[0]["req_num"] == 3 + assert prepare_calls[0]["plan"] == SpecDecodePlan(12, 7, 2, 3) + assert prepare_calls[1]["req_num"] == 2 + assert prepare_calls[1]["plan"] == SpecDecodePlan(8, 4, 2, 3) + + +@pytest.mark.parametrize( + ("batch_size0", "batch_size1", "dynamic_batch_size", "expected_batch_sizes"), + ( + (4, 8, 9, (4, 5)), + (8, 4, 10, (6, 4)), + ), +) +def test_dp_overlap_engine_moves_verify_rows_to_the_side_with_capacity( + batch_size0, + batch_size1, + dynamic_batch_size, + expected_batch_sizes, +): + dynamic_batch_sizes = [] + + class _CommonEngine: + def prepare_decode_model_input(self, **kwargs): + dynamic_batch_sizes.append(kwargs["plan"].dynamic_batch_size) + return kwargs["model_input"], None + + engine = DPOverlapSpecEngine.__new__(DPOverlapSpecEngine) + engine.common_engine = _CommonEngine() + engine.prepare_decode_model_inputs( + model_input0=SimpleNamespace(batch_size=batch_size0), + req_num0=2, + model_input1=SimpleNamespace(batch_size=batch_size1), + req_num1=2, + plan=SpecDecodePlan( + origin_batch_size=batch_size0 + batch_size1, + dynamic_batch_size=dynamic_batch_size, + draft_step=2, + pre_draft_step=3, + ), + ) + + assert tuple(dynamic_batch_sizes) == expected_batch_sizes + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_dp_overlap_decode_delegates_empty_layout_and_frees_proposal(monkeypatch): device = "cuda" @@ -327,14 +439,25 @@ def test_dp_overlap_decode_delegates_empty_layout_and_frees_proposal(monkeypatch calls = [] class _OverlapEngine: + def plan_decode(self, **kwargs): + calls.append("plan") + return SpecDecodePlan(0, 0, 2, 2) + + def prepare_decode_model_inputs(self, **kwargs): + calls.append("prepare") + return kwargs["model_input0"], None, kwargs["model_input1"], None + def propose_next_overlap(self, **kwargs): calls.append("propose") - assert kwargs["target_next_token_ids"].shape == (0,) - assert kwargs["target_next_token_ids"].dtype == torch.int64 - assert kwargs["real_verify_rows0"] == 0 - assert kwargs["real_verify_rows1"] == 0 + assert kwargs["target_next_token_ids0"].shape == (0,) + assert kwargs["target_next_token_ids0"].dtype == torch.int64 + assert kwargs["target_next_token_ids1"].shape == (0,) + assert kwargs["target_next_token_ids1"].dtype == torch.int64 + assert kwargs["real_request_num0"] == 0 + assert kwargs["real_request_num1"] == 0 assert kwargs["accept_len"].shape == (0,) assert kwargs["accept_len"].dtype == torch.int32 + assert kwargs["draft_step"] == 2 return SpecProposal(token_ids=torch.empty((0, 2), dtype=torch.int64, device=device)) backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) @@ -365,4 +488,12 @@ def propose_next_overlap(self, **kwargs): backend.decode_overlap_mtp(event_pack=event_pack, decode_reqs=[]) - assert calls == ["propose", "post_wait", "forward_wait", "free", "pre_post"] + assert calls == [ + "plan", + "prepare", + "propose", + "post_wait", + "forward_wait", + "free", + "pre_post", + ] diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 88ddc8dcfe..7a4ee8d889 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -1,14 +1,22 @@ from types import SimpleNamespace +import pytest import torch from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import DpOverlapEagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import ( + DpOverlapEagle3Proposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_no_att import ( + DpOverlapEagleNoAttProposer, +) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_with_att import ( DpOverlapEagleWithAttProposer, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import ( + Eagle3Proposer, +) class _DraftModel: @@ -34,8 +42,17 @@ def _microbatch_overlap_decode_cuda(self, input0, input1): self.decode_inputs.append((input0, input1)) return tuple( ModelOutput( - logits=torch.arange(model_input.batch_size, dtype=torch.float32).view(-1, 1), - mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), + logits=torch.arange( + model_input.batch_size, + dtype=torch.float32, + device=model_input.input_ids.device, + ).view(-1, 1), + mtp_collector=ModelMtpOutputCollector( + spec_hidden=torch.ones( + (model_input.batch_size, 2), + device=model_input.input_ids.device, + ) + ), ) for model_input in (input0, input1) ) @@ -44,18 +61,20 @@ def map_draft_vocab_to_main_vocab(self, token_ids): return token_ids -def _target_input(batch_size): +def _target_input(batch_size, b_mtp_index=None, device="cpu"): + if b_mtp_index is None: + b_mtp_index = torch.arange(batch_size, dtype=torch.int32, device=device) % 3 return SimpleNamespace( batch_size=batch_size, total_token_num=batch_size, - input_ids=torch.arange(batch_size, dtype=torch.int64), - b_seq_len=torch.arange(batch_size, dtype=torch.int32) + 4, - b_req_idx=torch.arange(batch_size, dtype=torch.int32), - b_mtp_index=torch.zeros(batch_size, dtype=torch.int32), - b_position_delta=torch.zeros(batch_size, dtype=torch.int32), - b_shared_seq_len=torch.zeros(batch_size, dtype=torch.int32), - b_shared_radix_node_id=torch.arange(batch_size, dtype=torch.int64), - mem_indexes=torch.arange(batch_size, dtype=torch.int32), + input_ids=torch.arange(batch_size, dtype=torch.int64, device=device), + b_seq_len=torch.arange(batch_size, dtype=torch.int32, device=device) + 4, + b_req_idx=torch.arange(batch_size, dtype=torch.int32, device=device), + b_mtp_index=b_mtp_index, + b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device=device), + b_shared_seq_len=torch.zeros(batch_size, dtype=torch.int32, device=device), + b_shared_radix_node_id=torch.arange(batch_size, dtype=torch.int64, device=device), + mem_indexes=torch.arange(batch_size, dtype=torch.int32, device=device), mem_indexes_cpu=torch.arange(batch_size, dtype=torch.int32), max_kv_seq_len=16, max_cache_len=16, @@ -64,7 +83,7 @@ def _target_input(batch_size): ) -def test_overlap_eagle_keeps_fixed_verify_layout(monkeypatch): +def test_overlap_eagle_supports_variable_verify_layout(monkeypatch): draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -82,42 +101,43 @@ def test_overlap_eagle_keeps_fixed_verify_layout(monkeypatch): "alloc_mem_indexes", lambda token_count: torch.arange(token_count, dtype=torch.int32), ) - model_input0 = _target_input(batch_size=6) + model_input0 = _target_input(batch_size=3) model_input1 = _target_input(batch_size=6) proposal = proposer.propose_next_overlap( target_model_input0=model_input0, target_model_output0=ModelOutput( - logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((3, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((3, 2))), ), + target_next_token_ids0=torch.arange(3, dtype=torch.int64), target_model_input1=model_input1, target_model_output1=ModelOutput( - logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) - ), - target_next_token_ids=torch.cat( - (torch.arange(3, dtype=torch.int64), torch.arange(10, 16, dtype=torch.int64)), - dim=0, + logits=torch.empty((6, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), ), - real_verify_rows0=3, - real_verify_rows1=6, + target_next_token_ids1=torch.arange(10, 16, dtype=torch.int64), + real_request_num0=1, + real_request_num1=2, accept_len=torch.tensor([2, 1, 3], dtype=torch.int32), draft_step=2, ) assert draft_model.extend_batch_sizes is None - assert draft_model.decode_batch_sizes == [(6, 6), (6, 6)] + assert draft_model.decode_batch_sizes == [(3, 6), (1, 2)] assert proposal.token_ids.shape == (3, 2) - expected_draft_tokens = torch.tensor([1, 0, 5]) - assert torch.equal(proposal.token_ids[:, 0], expected_draft_tokens) - assert torch.equal(proposal.token_ids[:, 1], expected_draft_tokens) + assert torch.equal(proposal.token_ids, torch.tensor([[1, 0], [0, 0], [5, 1]])) assert len(proposal.extra_mem_indexes_cpu) == 1 - assert torch.equal(proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, torch.arange(6, dtype=torch.int32)) + assert torch.equal( + proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, + torch.arange(3, dtype=torch.int32), + ) assert proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None - assert torch.equal(model_input0.mem_indexes, torch.tensor([2, 0, 1, 5, 99, 99], dtype=torch.int32)) - assert torch.equal(model_input1.mem_indexes, torch.tensor([2, 2, 4, 5, 3, 5], dtype=torch.int32)) + assert torch.equal(model_input0.mem_indexes, torch.arange(3, dtype=torch.int32)) + assert torch.equal(model_input1.mem_indexes, torch.arange(6, dtype=torch.int32)) -def test_overlap_eagle_builds_padded_inputs_for_empty_verify_rows(monkeypatch): +def test_overlap_eagle_supports_empty_verify_rows(monkeypatch): draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -137,23 +157,134 @@ def test_overlap_eagle_builds_padded_inputs_for_empty_verify_rows(monkeypatch): ) proposal = proposer.propose_next_overlap( - target_model_input0=_target_input(batch_size=6), + target_model_input0=_target_input(batch_size=0), target_model_output0=ModelOutput( - logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((0, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((0, 2))), ), - target_model_input1=_target_input(batch_size=6), + target_next_token_ids0=torch.empty((0,), dtype=torch.int64), + target_model_input1=_target_input(batch_size=0), target_model_output1=ModelOutput( - logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((0, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((0, 2))), ), - target_next_token_ids=torch.empty((0,), dtype=torch.int64), - real_verify_rows0=0, - real_verify_rows1=0, + target_next_token_ids1=torch.empty((0,), dtype=torch.int64), + real_request_num0=0, + real_request_num1=0, accept_len=torch.empty((0,), dtype=torch.int32), draft_step=2, ) assert proposal.token_ids.shape == (0, 2) - assert draft_model.decode_batch_sizes == [(6, 6), (6, 6)] + assert draft_model.decode_batch_sizes == [(0, 0), (0, 0)] + + +def test_overlap_eagle_returns_dynamic_schedule_scores(monkeypatch): + draft_model = _DraftModel() + backend = SimpleNamespace( + max_draft_step=2, + draft_models=[draft_model], + model=SimpleNamespace( + req_manager=SimpleNamespace( + mem_manager=SimpleNamespace(HOLD_TOKEN_MEMINDEX=99), + ) + ), + _gen_argmax_token_ids_and_prob=lambda output: ( + output.logits[:, 0].to(torch.int64), + output.logits[:, 0] / 10 + 0.5, + ), + ) + proposer = DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=True) + monkeypatch.setattr( + mtp_utils, + "alloc_mem_indexes", + lambda token_count: torch.arange(token_count, dtype=torch.int32), + ) + + proposal = proposer.propose_next_overlap( + target_model_input0=_target_input( + batch_size=2, + b_mtp_index=torch.tensor([0, 1], dtype=torch.int32), + ), + target_model_output0=ModelOutput( + logits=torch.empty((2, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((2, 2))), + ), + target_next_token_ids0=torch.arange(2, dtype=torch.int64), + target_model_input1=_target_input( + batch_size=3, + b_mtp_index=torch.tensor([0, 1, 0], dtype=torch.int32), + ), + target_model_output1=ModelOutput( + logits=torch.empty((3, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((3, 2))), + ), + target_next_token_ids1=torch.arange(2, 5, dtype=torch.int64), + real_request_num0=1, + real_request_num1=2, + accept_len=torch.tensor([2, 1, 1], dtype=torch.int32), + draft_step=2, + ) + + assert torch.equal(proposal.token_ids, torch.tensor([[1, 0], [0, 0], [2, 1]])) + assert torch.allclose( + proposal.schedule_scores, + torch.tensor([[0.6, 0.5], [0.5, 0.5], [0.7, 0.6]]), + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_overlap_eagle_no_att_supports_dynamic_draft_step(): + device = "cuda" + draft_model = _DraftModel() + backend = SimpleNamespace( + max_draft_step=3, + draft_models=[draft_model], + _gen_argmax_token_ids_and_prob=lambda output: ( + output.logits[:, 0].to(torch.int64), + output.logits[:, 0] / 10 + 0.5, + ), + ) + proposer = DpOverlapEagleNoAttProposer(backend=backend, enable_dynmaic_mtp=True) + model_input0 = _target_input( + batch_size=2, + b_mtp_index=torch.tensor([0, 1], dtype=torch.int32, device=device), + device=device, + ) + model_input1 = _target_input( + batch_size=3, + b_mtp_index=torch.tensor([0, 1, 0], dtype=torch.int32, device=device), + device=device, + ) + + proposal = proposer.propose_next_overlap( + target_model_input0=model_input0, + target_model_output0=ModelOutput( + logits=torch.empty((2, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((2, 2), device=device)), + ), + target_next_token_ids0=torch.arange(2, dtype=torch.int64, device=device), + target_model_input1=model_input1, + target_model_output1=ModelOutput( + logits=torch.empty((3, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((3, 2), device=device)), + ), + target_next_token_ids1=torch.arange(2, 5, dtype=torch.int64, device=device), + real_request_num0=1, + real_request_num1=2, + accept_len=torch.tensor([2, 1, 1], dtype=torch.int32, device=device), + draft_step=2, + ) + + assert proposal.token_ids.tolist() == [[0, 0], [0, 0], [1, 1]] + assert torch.allclose( + proposal.schedule_scores, + torch.tensor( + [[0.5, 0.5], [0.5, 0.5], [0.6, 0.6]], + device=device, + ), + ) + assert draft_model.decode_batch_sizes == [(1, 2), (1, 2)] def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): @@ -174,38 +305,41 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): "alloc_mem_indexes", lambda token_count: torch.arange(token_count, dtype=torch.int32), ) - model_input0 = _target_input(batch_size=6) + model_input0 = _target_input(batch_size=3) model_input1 = _target_input(batch_size=6) proposal = proposer.propose_next_overlap( target_model_input0=model_input0, target_model_output0=ModelOutput( - logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((3, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((3, 2))), ), + target_next_token_ids0=torch.arange(3, dtype=torch.int64), target_model_input1=model_input1, target_model_output1=ModelOutput( - logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))) + logits=torch.empty((6, 1)), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), ), - target_next_token_ids=torch.cat( - (torch.arange(3, dtype=torch.int64), torch.arange(10, 16, dtype=torch.int64)), - dim=0, - ), - real_verify_rows0=3, - real_verify_rows1=6, + target_next_token_ids1=torch.arange(10, 16, dtype=torch.int64), + real_request_num0=1, + real_request_num1=2, accept_len=torch.tensor([2, 1, 3], dtype=torch.int32), draft_step=2, ) assert draft_model.extend_inputs is None assert len(draft_model.decode_inputs) == 2 - assert draft_model.decode_inputs[0][0] is model_input0 - assert draft_model.decode_inputs[0][1] is model_input1 + assert draft_model.decode_inputs[0][0] is not model_input0 + assert draft_model.decode_inputs[0][1] is not model_input1 assert draft_model.extend_batch_sizes is None - assert draft_model.decode_batch_sizes == [(6, 6), (2, 2)] + assert draft_model.decode_batch_sizes == [(3, 6), (1, 2)] assert proposal.token_ids.shape == (3, 2) assert torch.equal(proposal.token_ids, torch.tensor([[1, 0], [0, 0], [5, 1]])) assert len(proposal.extra_mem_indexes_cpu) == 1 - assert torch.equal(proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, torch.arange(3, dtype=torch.int32)) + assert torch.equal( + proposal.extra_mem_indexes_cpu[0].mem_indexes_cpu, + torch.arange(3, dtype=torch.int32), + ) assert proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None @@ -214,7 +348,10 @@ def test_eagle3_maps_draft_token_ids_in_proposer(): proposer.backend = SimpleNamespace( draft_models=[SimpleNamespace(map_draft_vocab_to_main_vocab=lambda token_ids: token_ids + 100)], _gen_argmax_token_ids=lambda _: torch.tensor([1, 2]), - _gen_argmax_token_ids_and_prob=lambda _: (torch.tensor([3, 4]), torch.tensor([0.8, 0.7])), + _gen_argmax_token_ids_and_prob=lambda _: ( + torch.tensor([3, 4]), + torch.tensor([0.8, 0.7]), + ), ) token_ids = proposer._gen_argmax_token_ids(ModelOutput(logits=torch.empty(0))) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 3f511ab0ac..54c7a0006f 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -5,16 +5,15 @@ import torch from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_planner import ( - BaseDpOverlapPlanner, - FixedDpOverlapPlanner, - build_dp_overlap_planner, -) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers import ( build_dp_overlap_spec_proposer, ) -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import BaseDpOverlapProposer -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import DpOverlapEagle3Proposer +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( + BaseDpOverlapProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import ( + DpOverlapEagle3Proposer, +) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_no_att import ( DpOverlapEagleNoAttProposer, ) @@ -34,26 +33,44 @@ LightSpecPlanner, SpecDecodePlan, ) -from lightllm.server.router.model_infer.mtp_speculative.planner.base import _InferCostMsTable -from lightllm.server.router.model_infer.mtp_speculative.proposers import build_spec_proposer +from lightllm.server.router.model_infer.mtp_speculative.planner.base import ( + _InferCostMsTable, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers import ( + build_spec_proposer, +) from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( BaseSpecProposer, MtpMemIndexesToFree, SpecProposal, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import DFlashProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import DSparkProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import Eagle3Proposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import EagleNoAttProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.dflash import ( + DFlashProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.dspark import ( + DSparkProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import ( + Eagle3Proposer, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_no_att import ( + EagleNoAttProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import ( + EagleWithAttProposer, +) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( DFlashSpecProposal, DSparkSpecProposal, EagleSpecProposal, VanillaSpecProposal, ) -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import VanillaNoAttProposer -from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import VanillaWithAttProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_no_att import ( + VanillaNoAttProposer, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.vanilla_with_att import ( + VanillaWithAttProposer, +) from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils @@ -61,10 +78,11 @@ def build_lightspec_planner( max_draft_step: int = 3, spec_mode: str = "vanilla_with_att", block_size: int = 3, + enable_decode_microbatch_overlap: bool = False, ): backend = SimpleNamespace( args=SimpleNamespace(dp=1), - enable_decode_microbatch_overlap=False, + enable_decode_microbatch_overlap=enable_decode_microbatch_overlap, max_draft_step=max_draft_step, model=SimpleNamespace(graph=None), draft_models=[SimpleNamespace(block_size=block_size, graph=None)], @@ -258,14 +276,6 @@ def test_scatter_mtp_next_tokens_ignores_empty_schedule_scores(monkeypatch): assert scatter_args["schedule_scores"] is None -def test_dp_overlap_planner_returns_fixed_backend_draft_step(): - planner = build_dp_overlap_planner(backend=SimpleNamespace(max_draft_step=5)) - - assert type(planner) is FixedDpOverlapPlanner - assert isinstance(planner, BaseDpOverlapPlanner) - assert planner.get_draft_step() == 5 - - def test_infer_cost_candidates_include_feasible_boundaries(): costs = _InferCostMsTable() costs.update(batch_size=4, infer_cost_ms=1.0) @@ -325,6 +335,28 @@ def test_dynamic_planner_registers_cuda_graph_costs_from_backend(): assert planner.draft_infer_costs.estimate(4) == 0.3 +def test_overlap_lightspec_consumes_combined_cuda_graph_batch_size(): + backend = SimpleNamespace( + args=SimpleNamespace(dp=1), + enable_decode_microbatch_overlap=True, + max_draft_step=3, + model=SimpleNamespace(graph=SimpleNamespace(infer_cost_ms_by_batch_size={8: 1.2})), + draft_models=[ + SimpleNamespace( + block_size=3, + graph=SimpleNamespace(infer_cost_ms_by_batch_size={8: 0.3}), + ) + ], + ) + + planner = LightSpecPlanner(spec_mode="vanilla_with_att", backend=backend) + + assert planner.target_infer_costs.estimate(7) == 1.2 + assert planner.target_infer_costs.estimate(8) == 1.2 + assert planner.draft_infer_costs.estimate(7) == 0.3 + assert planner.draft_infer_costs.estimate(8) == 0.3 + + def test_each_mode_proposer_inherits_its_expected_implementation_base(): proposer_types = ( VanillaWithAttProposer, @@ -511,6 +543,44 @@ def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): assert plan.draft_step == 1 +def test_lightspec_overlap_prices_combined_batch_size(): + planner = build_lightspec_planner( + spec_mode="eagle_with_att", + enable_decode_microbatch_overlap=True, + ) + planner.target_infer_costs.update(batch_size=4, infer_cost_ms=1.0) + planner.target_infer_costs.update(batch_size=8, infer_cost_ms=3.0) + planner.draft_infer_costs.update(batch_size=4, infer_cost_ms=0.25) + planner.draft_infer_costs.update(batch_size=8, infer_cost_ms=1.0) + + assert planner.target_infer_costs.estimate(7) == 3.0 + assert planner.target_infer_costs.get_batch_size_keys_between(4, 8) == [4, 8] + assert planner._get_cost_ms(req_num=4, dynamic_batch_size=8, draft_step=1) == 0.5 + + +def test_lightspec_overlap_selects_dynamic_draft_step(): + planner = build_lightspec_planner( + spec_mode="eagle_with_att", + enable_decode_microbatch_overlap=True, + ) + for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 10.0)): + planner.target_infer_costs.update(batch_size=batch_size, infer_cost_ms=target_cost) + for batch_size, draft_cost in ((2, 0.1), (4, 0.1), (8, 0.2)): + planner.draft_infer_costs.update(batch_size=batch_size, infer_cost_ms=draft_cost) + planner._update_verified_batch( + accept_lengths=[4, 4], + req_num=2, + dynamic_batch_size=8, + verified_draft_step=3, + ) + + plan = planner.plan(decode_reqs=build_decode_reqs(2), origin_batch_size=8) + + assert plan.dynamic_batch_size == 4 + assert plan.pre_draft_step == 3 + assert plan.draft_step == 1 + + def test_lightspec_keeps_vanilla_attention_chained_depth_fixed(): planner = build_lightspec_planner(spec_mode="vanilla_with_att") for batch_size, target_cost in ((2, 1.0), (4, 1.1), (8, 10.0)): @@ -575,11 +645,17 @@ def test_engine_lets_planner_count_requests_with_a_previous_proposal(): mixed_plan = engine.plan_decode( model_input=model_input, - decode_reqs=[SimpleNamespace(cur_output_len=1), SimpleNamespace(cur_output_len=2)], + decode_reqs=[ + SimpleNamespace(cur_output_len=1), + SimpleNamespace(cur_output_len=2), + ], ) ready_plan = engine.plan_decode( model_input=model_input, - decode_reqs=[SimpleNamespace(cur_output_len=2), SimpleNamespace(cur_output_len=2)], + decode_reqs=[ + SimpleNamespace(cur_output_len=2), + SimpleNamespace(cur_output_len=2), + ], ) assert mixed_plan.dynamic_batch_size == 5 @@ -591,7 +667,9 @@ def test_engine_lets_planner_count_requests_with_a_previous_proposal(): def test_dynamic_prepare_keeps_prefix_mem_indexes_and_frees_unused_tail(monkeypatch): from lightllm.common.basemodel.triton_kernel import dynamic_mtp_utils from lightllm.server.router.model_infer import infer_batch as infer_batch_module - from lightllm.server.router.model_infer.mtp_speculative import engine as engine_module + from lightllm.server.router.model_infer.mtp_speculative import ( + engine as engine_module, + ) freed = [] monkeypatch.setattr( diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py index cc17aba811..1a40d33448 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -4,6 +4,9 @@ import torch from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_no_att import ( + DpOverlapVanillaNoAttProposer, +) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.vanilla_with_att import ( DpOverlapVanillaWithAttProposer, ) @@ -33,6 +36,31 @@ def _microbatch_overlap_decode_cuda(self, input0, input1): ) +def test_dp_vanilla_no_att_supports_zero_dynamic_draft_step(): + proposer = DpOverlapVanillaNoAttProposer( + backend=SimpleNamespace(draft_models=[]), + enable_dynmaic_mtp=True, + ) + target_next_token_ids0 = torch.tensor([10], dtype=torch.int64) + target_next_token_ids1 = torch.tensor([11], dtype=torch.int64) + + proposal = proposer.propose_next_overlap( + target_model_input0=None, + target_model_output0=None, + target_next_token_ids0=target_next_token_ids0, + target_model_input1=None, + target_model_output1=None, + target_next_token_ids1=target_next_token_ids1, + real_request_num0=1, + real_request_num1=1, + accept_len=torch.ones((2,), dtype=torch.int32), + draft_step=0, + ) + + assert proposal.token_ids.shape == (2, 0) + assert proposal.schedule_scores.shape == (2, 0) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_dp_vanilla_proposer_owns_overlap_decode(): device = "cuda" @@ -43,8 +71,17 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) - model_input0 = SimpleNamespace(batch_size=6, b_req_idx=torch.arange(6, dtype=torch.int32, device=device)) - model_input1 = SimpleNamespace(batch_size=6, b_req_idx=torch.arange(6, dtype=torch.int32, device=device)) + b_mtp_index = torch.arange(3, dtype=torch.int32, device=device) + model_input0 = SimpleNamespace( + batch_size=3, + b_req_idx=torch.arange(3, dtype=torch.int32, device=device), + b_mtp_index=b_mtp_index, + ) + model_input1 = SimpleNamespace( + batch_size=3, + b_req_idx=torch.arange(3, dtype=torch.int32, device=device), + b_mtp_index=b_mtp_index, + ) proposal = proposer.propose_next_overlap( target_model_input0=model_input0, @@ -52,14 +89,15 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): logits=torch.empty((6, 1), device=device), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), ), + target_next_token_ids0=torch.tensor([10, 11, 0], dtype=torch.int64, device=device), target_model_input1=model_input1, target_model_output1=ModelOutput( logits=torch.empty((6, 1), device=device), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), ), - target_next_token_ids=torch.tensor([10, 11, 0, 20, 21, 22], dtype=torch.int64, device=device), - real_verify_rows0=3, - real_verify_rows1=3, + target_next_token_ids1=torch.tensor([20, 21, 22], dtype=torch.int64, device=device), + real_request_num0=1, + real_request_num1=1, accept_len=torch.tensor([2, 1], dtype=torch.int32, device=device), draft_step=2, ) @@ -69,8 +107,8 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): [0, 0], ] assert proposal.extra_mem_indexes_cpu == [] - assert draft_models[0].decode_batch_sizes == [(6, 6)] - assert draft_models[1].decode_batch_sizes == [(6, 6)] + assert draft_models[0].decode_batch_sizes == [(3, 3)] + assert draft_models[1].decode_batch_sizes == [(3, 3)] @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") @@ -83,29 +121,31 @@ def test_dp_vanilla_proposer_builds_padded_inputs_for_empty_verify_rows(): _gen_argmax_token_ids=lambda output: output.logits[:, 0].to(torch.int64), ) proposer = DpOverlapVanillaWithAttProposer(backend=backend, enable_dynmaic_mtp=False) - model_input0 = SimpleNamespace(batch_size=6, b_req_idx=torch.arange(6, dtype=torch.int32, device=device)) - model_input1 = SimpleNamespace(batch_size=6, b_req_idx=torch.arange(6, dtype=torch.int32, device=device)) + empty_i32 = torch.empty((0,), dtype=torch.int32, device=device) + model_input0 = SimpleNamespace(batch_size=0, b_req_idx=empty_i32, b_mtp_index=empty_i32) + model_input1 = SimpleNamespace(batch_size=0, b_req_idx=empty_i32, b_mtp_index=empty_i32) model_output0 = ModelOutput( - logits=torch.empty((6, 1), device=device), - mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), + logits=torch.empty((0, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((0, 2), device=device)), ) model_output1 = ModelOutput( - logits=torch.empty((6, 1), device=device), - mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), + logits=torch.empty((0, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((0, 2), device=device)), ) proposal = proposer.propose_next_overlap( target_model_input0=model_input0, target_model_output0=model_output0, + target_next_token_ids0=torch.empty((0,), dtype=torch.int64, device=device), target_model_input1=model_input1, target_model_output1=model_output1, - target_next_token_ids=torch.empty((0,), dtype=torch.int64, device=device), - real_verify_rows0=0, - real_verify_rows1=0, + target_next_token_ids1=torch.empty((0,), dtype=torch.int64, device=device), + real_request_num0=0, + real_request_num1=0, accept_len=torch.empty((0,), dtype=torch.int32, device=device), draft_step=2, ) assert proposal.token_ids.shape == (0, 2) - assert draft_models[0].decode_batch_sizes == [(6, 6)] - assert draft_models[1].decode_batch_sizes == [(6, 6)] + assert draft_models[0].decode_batch_sizes == [(0, 0)] + assert draft_models[1].decode_batch_sizes == [(0, 0)] From f0f38e6ce0e68f56bc70263c6fe00c12ff622b68 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sat, 22 Aug 2026 03:43:34 +0000 Subject: [PATCH 090/103] refactor dp overlap eagle proposers --- .../dp_overlap_proposers/eagle3.py | 68 +-- .../dp_overlap_proposers/eagle_no_att.py | 142 ++++-- .../dp_overlap_proposers/eagle_utils.py | 439 ------------------ .../dp_overlap_proposers/eagle_with_att.py | 255 ++++++++-- .../dp_overlap_proposers/utils.py | 17 + .../dp_overlap_proposers/vanilla_no_att.py | 131 +++++- .../dp_overlap_proposers/vanilla_with_att.py | 18 +- .../mtp_speculative/test_eagle_overlap.py | 43 ++ .../mtp_speculative/test_planner.py | 2 +- 9 files changed, 539 insertions(+), 576 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py create mode 100644 lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/utils.py diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py index 2f58a95794..6594182db7 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -1,67 +1,21 @@ import torch -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( - BaseDpOverlapProposer, -) -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( - fill_dp_eagle_draft_model_kv_state_overlap, - propose_next_dp_eagle_autoregressive_overlap, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( - EagleSpecProposal, +from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_with_att import ( + DpOverlapEagleWithAttProposer, ) -class DpOverlapEagle3Proposer(BaseDpOverlapProposer): - """DP ``eagle3`` proposer。""" +class DpOverlapEagle3Proposer(DpOverlapEagleWithAttProposer): + """复用 DP EAGLE With-Att KV 流程并执行 draft-to-target 词表映射。""" def _map_draft_token_ids(self, draft_token_ids: torch.Tensor) -> torch.Tensor: return self.backend.draft_models[0].map_draft_vocab_to_main_vocab(draft_token_ids) - def fill_draft_model_kv_state_overlap( - self, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, - ) -> None: - fill_dp_eagle_draft_model_kv_state_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - ) + def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: + draft_token_ids = super()._gen_argmax_token_ids(model_output) + return self._map_draft_token_ids(draft_token_ids) - def propose_next_overlap( - self, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, - draft_step: int, - ) -> EagleSpecProposal: - return propose_next_dp_eagle_autoregressive_overlap( - proposer=self, - target_model_input0=target_model_input0, - target_model_output0=target_model_output0, - target_next_token_ids0=target_next_token_ids0, - target_model_input1=target_model_input1, - target_model_output1=target_model_output1, - target_next_token_ids1=target_next_token_ids1, - real_request_num0=real_request_num0, - real_request_num1=real_request_num1, - accept_len=accept_len, - draft_step=draft_step, - map_draft_token_ids=self._map_draft_token_ids, - ) + def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): + draft_token_ids, draft_token_probs = super()._gen_argmax_token_ids_and_prob(model_output) + return self._map_draft_token_ids(draft_token_ids), draft_token_probs diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 546d25b765..2843ba496b 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -1,12 +1,14 @@ +import copy + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( BaseDpOverlapProposer, ) -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( - fill_dp_eagle_draft_model_kv_state_overlap, - propose_next_dp_no_att_overlap, +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.utils import ( + get_dp_overlap_req_start_rows, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( EagleSpecProposal, @@ -25,15 +27,7 @@ def fill_draft_model_kv_state_overlap( target_model_output1: ModelOutput, target_next_token_ids1: torch.Tensor, ) -> None: - fill_dp_eagle_draft_model_kv_state_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - ) + pass def propose_next_overlap( self, @@ -48,18 +42,114 @@ def propose_next_overlap( accept_len: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: - return propose_next_dp_no_att_overlap( - proposer=self, - target_model_input0=target_model_input0, - target_model_output0=target_model_output0, - target_next_token_ids0=target_next_token_ids0, - target_model_input1=target_model_input1, - target_model_output1=target_model_output1, - target_next_token_ids1=target_next_token_ids1, - real_request_num0=real_request_num0, - real_request_num1=real_request_num1, - accept_len=accept_len, - draft_step=draft_step, - get_draft_model=lambda _: self.backend.draft_models[0], - map_draft_token_ids=lambda token_ids: token_ids, + req_num_by_batch = (int(real_request_num0), int(real_request_num1)) + req_num = sum(req_num_by_batch) + assert accept_len.shape == (req_num,) + + proposal_token_ids = target_next_token_ids0.new_empty((req_num, draft_step)) + schedule_scores = ( + torch.empty( + (req_num, draft_step), + dtype=torch.float32, + device=target_next_token_ids0.device, + ) + if self.enable_dynmaic_mtp + else None + ) + + if draft_step == 0: + return EagleSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, + ) + + assert target_next_token_ids0.shape == (target_model_input0.batch_size,) + assert target_next_token_ids1.shape == (target_model_input1.batch_size,) + accept_len_by_batch = ( + accept_len[: req_num_by_batch[0]], + accept_len[req_num_by_batch[0] :], + ) + b_req_mtp_start_loc = ( + get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input0.b_mtp_index, + req_num=req_num_by_batch[0], + ), + get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input1.b_mtp_index, + req_num=req_num_by_batch[1], + ), + ) + + model_inputs = (target_model_input0, target_model_input1) + model_outputs = (target_model_output0, target_model_output1) + target_next_token_ids_by_batch = (target_next_token_ids0, target_next_token_ids1) + draft_inputs = [] + draft_token_ids_by_batch = [] + draft_hiddens_by_batch = [] + for model_input, model_output, token_ids, req_mtp_start_loc, batch_accept_len, batch_req_num in zip( + model_inputs, + model_outputs, + target_next_token_ids_by_batch, + b_req_mtp_start_loc, + accept_len_by_batch, + req_num_by_batch, + ): + selected_rows = select_accepted_tail_rows( + b_req_mtp_start_loc=req_mtp_start_loc, + accept_len=batch_accept_len, + input_ids=token_ids, + hidden=model_output.mtp_collector.spec_hidden, + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, + b_seq_len=model_input.b_seq_len, + mem_indexes=model_input.mem_indexes, + b_shared_seq_len=model_input.b_shared_seq_len, + b_shared_radix_node_id=model_input.b_shared_radix_node_id, + b_position_delta=model_input.b_position_delta, + ) + draft_input = copy.copy(model_input) + draft_input.batch_size = batch_req_num + draft_input.input_ids = selected_rows.input_ids + draft_input.b_req_idx = selected_rows.b_req_idx + draft_input.b_mtp_index = selected_rows.b_mtp_index + draft_input.b_seq_len = selected_rows.b_seq_len + draft_input.mem_indexes = selected_rows.mem_indexes + draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len + draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id + draft_input.b_position_delta = selected_rows.b_position_delta + draft_input.mem_indexes_cpu = None + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(batch_req_num)] + draft_inputs.append(draft_input) + draft_token_ids_by_batch.append(selected_rows.input_ids) + draft_hiddens_by_batch.append(selected_rows.hidden) + + proposal_row_offsets = (0, req_num_by_batch[0]) + draft_model = self.backend.draft_models[0] + for step in range(draft_step): + for batch_index, draft_input in enumerate(draft_inputs): + draft_input.input_ids = draft_token_ids_by_batch[batch_index] + draft_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] + + draft_outputs = draft_model._microbatch_overlap_decode_cuda(*draft_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + draft_token_probs = draft_token_probs.float() + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + draft_token_ids_by_batch[batch_index] = draft_token_ids + draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden + + batch_req_num = req_num_by_batch[batch_index] + proposal_row_start = proposal_row_offsets[batch_index] + proposal_row_end = proposal_row_start + batch_req_num + proposal_token_ids[proposal_row_start:proposal_row_end, step] = draft_token_ids + if schedule_scores is not None: + schedule_scores[proposal_row_start:proposal_row_end, step] = draft_token_probs + + return EagleSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py deleted file mode 100644 index 1b818b5f3c..0000000000 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_utils.py +++ /dev/null @@ -1,439 +0,0 @@ -"""DP overlap EAGLE proposer 共享辅助函数。""" - -from __future__ import annotations - -import copy -from typing import Callable - -import torch - -from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput -from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import ( - gen_mtp_new_input_ids, -) -from lightllm.common.basemodel.triton_kernel.mtp_utils import gen_b_req_mtp_start_loc -from lightllm.common.basemodel.triton_kernel.select_mtp_rows import ( - select_accepted_tail_rows, -) -from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( - BaseDpOverlapProposer, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( - MtpMemIndexesToFree, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( - EagleSpecProposal, -) -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager - - -def _prepare_eagle_prefill_inputs( - model_input: ModelInput, - b_next_token_ids: torch.Tensor, - mtp_draft_input_hiddens: torch.Tensor, -) -> None: - """构造 DP-overlap EAGLE draft model 的 prefill 输入。""" - - model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( - key="dp_overlap_eagle_prefill_b_is_decode_req", - shape=model_input.b_req_idx.shape, - fill_value=False, - dtype=torch.bool, - ) - model_input.input_ids = gen_mtp_new_input_ids( - input_ids=model_input.input_ids, - b_next_token_ids=b_next_token_ids, - b_seq_len=model_input.b_seq_len, - b_ready_cache_len=model_input.b_ready_cache_len, - ) - model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens - - -def generate_eagle_tokens( - proposer: BaseDpOverlapProposer, - model_output: ModelOutput, - map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> tuple[torch.Tensor, torch.Tensor | None]: - if proposer.enable_dynmaic_mtp: - draft_token_ids, draft_token_probs = proposer.backend._gen_argmax_token_ids_and_prob(model_output) - return map_draft_token_ids(draft_token_ids), draft_token_probs.float() - return ( - map_draft_token_ids(proposer.backend._gen_argmax_token_ids(model_output)), - None, - ) - - -def prepare_eagle_verify_decode_input( - model_input: ModelInput, - input_ids: torch.Tensor, - target_hidden: torch.Tensor, -) -> None: - """复用 target verify 布局构造 DP-overlap drafter 输入。""" - - assert not model_input.is_prefill - model_input.input_ids = input_ids - model_input.mtp_draft_input_hiddens = target_hidden - - -def fill_dp_eagle_draft_model_kv_state_overlap( - proposer: BaseDpOverlapProposer, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, -) -> None: - """使用两个 target prefill microbatch 初始化 EAGLE draft state。""" - - _prepare_eagle_prefill_inputs( - model_input=target_model_input0, - b_next_token_ids=target_next_token_ids0, - mtp_draft_input_hiddens=target_model_output0.mtp_collector.spec_hidden, - ) - _prepare_eagle_prefill_inputs( - model_input=target_model_input1, - b_next_token_ids=target_next_token_ids1, - mtp_draft_input_hiddens=target_model_output1.mtp_collector.spec_hidden, - ) - proposer.backend.draft_models[0]._microbatch_overlap_prefill_cuda(target_model_input0, target_model_input1) - - -def get_dp_overlap_req_start_rows(b_mtp_index: torch.Tensor, request_num: int) -> torch.Tensor: - """根据动态 verify 布局恢复请求起始行。""" - - request_num = int(request_num) - if request_num == 0: - assert b_mtp_index.numel() == 0 - return torch.empty((0,), dtype=torch.int32, device=b_mtp_index.device) - if b_mtp_index.is_cuda: - return gen_b_req_mtp_start_loc( - b_mtp_index=b_mtp_index, - num_reqs=request_num, - ) - req_start_rows = torch.nonzero(b_mtp_index == 0, as_tuple=False).flatten().to(dtype=torch.int32) - assert req_start_rows.shape == (request_num,) - return req_start_rows - - -def prepare_dp_eagle_overlap_decode_inputs( - target_model_input0: ModelInput, - target_model_input1: ModelInput, - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - """恢复两个动态 verify microbatch 的请求起始行。""" - - assert accept_len.shape == (real_request_num0 + real_request_num1,) - accept_len0 = accept_len[:real_request_num0] - accept_len1 = accept_len[real_request_num0:] - req_start_rows0 = get_dp_overlap_req_start_rows( - b_mtp_index=target_model_input0.b_mtp_index, - request_num=real_request_num0, - ) - req_start_rows1 = get_dp_overlap_req_start_rows( - b_mtp_index=target_model_input1.b_mtp_index, - request_num=real_request_num1, - ) - return ( - accept_len0, - accept_len1, - req_start_rows0, - req_start_rows1, - ) - - -def propose_next_dp_eagle_autoregressive_overlap( - proposer: BaseDpOverlapProposer, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, - draft_step: int, - map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> EagleSpecProposal: - """运行 DP EAGLE extend 后接单 token overlap decode 的 proposal 流程。""" - - assert target_next_token_ids0.shape == (target_model_input0.batch_size,) - assert target_next_token_ids1.shape == (target_model_input1.batch_size,) - (accept_len0, accept_len1, req_start_rows0, req_start_rows1,) = prepare_dp_eagle_overlap_decode_inputs( - target_model_input0=target_model_input0, - target_model_input1=target_model_input1, - real_request_num0=real_request_num0, - real_request_num1=real_request_num1, - accept_len=accept_len, - ) - - model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] - model_outputs = (target_model_output0, target_model_output1) - target_next_token_ids_by_batch = (target_next_token_ids0, target_next_token_ids1) - accept_lens_by_batch = (accept_len0, accept_len1) - req_start_rows_by_batch = (req_start_rows0, req_start_rows1) - position_deltas_by_batch = tuple(model_input.b_position_delta for model_input in model_inputs) - max_kv_seq_lens_by_batch = tuple(model_input.max_kv_seq_len for model_input in model_inputs) - - real_request_counts = (int(real_request_num0), int(real_request_num1)) - accepted_tail_rows_by_batch = [] - for model_input, model_output, token_ids, req_start_rows, accept_len in zip( - model_inputs, - model_outputs, - target_next_token_ids_by_batch, - req_start_rows_by_batch, - accept_lens_by_batch, - ): - accepted_tail_rows = (req_start_rows + accept_len - 1).long() - accepted_tail_rows_by_batch.append(accepted_tail_rows) - prepare_eagle_verify_decode_input( - model_input=model_input, - input_ids=token_ids, - target_hidden=model_output.mtp_collector.spec_hidden, - ) - - total_real_request_count = sum(real_request_counts) - proposal_token_ids = target_next_token_ids0.new_empty((total_real_request_count, draft_step)) - proposal_schedule_scores = ( - torch.empty( - (total_real_request_count, draft_step), - dtype=torch.float32, - device=target_next_token_ids0.device, - ) - if proposer.enable_dynmaic_mtp - else None - ) - - draft_model = proposer.backend.draft_models[0] - extend_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) - - draft_token_ids_by_batch = [] - draft_hiddens_by_batch = [] - draft_seq_lens_by_batch = [] - draft_req_indices_by_batch = [] - draft_shared_seq_lens_by_batch = [] - draft_shared_radix_node_ids_by_batch = [] - proposal_row_offsets = (0, real_request_counts[0]) - for batch_index, (model_input, extend_output, accepted_tail_rows, real_request_count,) in enumerate( - zip( - model_inputs, - extend_outputs, - accepted_tail_rows_by_batch, - real_request_counts, - ) - ): - accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) - draft_token_ids, draft_token_probs = generate_eagle_tokens( - proposer, - accepted_tail_output, - map_draft_token_ids, - ) - draft_token_ids_by_batch.append(draft_token_ids) - draft_hiddens_by_batch.append(extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows)) - draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) - draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) - draft_shared_seq_lens_by_batch.append(model_input.b_shared_seq_len.index_select(0, accepted_tail_rows)) - draft_shared_radix_node_ids_by_batch.append( - model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) - ) - proposal_row_start = proposal_row_offsets[batch_index] - proposal_token_ids[proposal_row_start : proposal_row_start + real_request_count, 0] = draft_token_ids[ - :real_request_count - ] - if proposal_schedule_scores is not None: - proposal_schedule_scores[ - proposal_row_start : proposal_row_start + real_request_count, 0 - ] = draft_token_probs[:real_request_count] - - if draft_step == 1: - return EagleSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=proposal_schedule_scores, - ) - - for batch_index, model_input in enumerate(model_inputs): - model_input.is_prefill = False - model_input.batch_size = real_request_counts[batch_index] - model_input.b_req_idx = draft_req_indices_by_batch[batch_index] - model_input.b_mtp_index = torch.zeros_like(model_input.b_req_idx) - model_input.b_seq_len = draft_seq_lens_by_batch[batch_index] - assert position_deltas_by_batch[batch_index] is not None - model_input.b_position_delta = position_deltas_by_batch[batch_index].index_select( - 0, accepted_tail_rows_by_batch[batch_index] - ) - model_input.b_shared_seq_len = draft_shared_seq_lens_by_batch[batch_index] - model_input.b_shared_radix_node_id = draft_shared_radix_node_ids_by_batch[batch_index] - if len(model_input.multimodal_params) != model_input.batch_size: - empty_multimodal_params = {"images": [], "audios": []} - model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size - - extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(total_real_request_count * (draft_step - 1)) - extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids0.device, non_blocking=True) - - for step in range(1, draft_step): - mem_start = (step - 1) * total_real_request_count - step_mem_indexes = extra_mem_indexes[mem_start : mem_start + total_real_request_count] - real_mem_start = 0 - for batch_index, model_input in enumerate(model_inputs): - real_request_count = real_request_counts[batch_index] - real_mem_indexes = step_mem_indexes[real_mem_start : real_mem_start + real_request_count] - real_mem_start += real_request_count - model_input.input_ids = draft_token_ids_by_batch[batch_index] - model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - model_input.mem_indexes = real_mem_indexes - model_input.max_kv_seq_len = max_kv_seq_lens_by_batch[batch_index] + step - model_input.total_token_num = model_input.batch_size * model_input.max_kv_seq_len - - draft_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) - for batch_index, draft_output in enumerate(draft_outputs): - draft_token_ids, draft_token_probs = generate_eagle_tokens( - proposer, - draft_output, - map_draft_token_ids, - ) - draft_token_ids_by_batch[batch_index] = draft_token_ids - draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden - draft_seq_lens_by_batch[batch_index].add_(1) - real_request_count = real_request_counts[batch_index] - proposal_row_start = proposal_row_offsets[batch_index] - proposal_token_ids[proposal_row_start : proposal_row_start + real_request_count, step] = draft_token_ids[ - :real_request_count - ] - if proposal_schedule_scores is not None: - proposal_schedule_scores[ - proposal_row_start : proposal_row_start + real_request_count, - step, - ] = draft_token_probs[:real_request_count] - - return EagleSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], - schedule_scores=proposal_schedule_scores, - ) - - -def propose_next_dp_no_att_overlap( - proposer: BaseDpOverlapProposer, - target_model_input0: ModelInput, - target_model_output0: ModelOutput, - target_next_token_ids0: torch.Tensor, - target_model_input1: ModelInput, - target_model_output1: ModelOutput, - target_next_token_ids1: torch.Tensor, - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, - draft_step: int, - get_draft_model: Callable[[int], object], - map_draft_token_ids: Callable[[torch.Tensor], torch.Tensor], -) -> EagleSpecProposal: - """将动态 verify 布局压缩为每请求一行,再执行双 microbatch No-Att draft。""" - - real_request_counts = (int(real_request_num0), int(real_request_num1)) - total_request_num = sum(real_request_counts) - proposal_token_ids = target_next_token_ids0.new_empty((total_request_num, draft_step)) - proposal_schedule_scores = ( - torch.empty( - (total_request_num, draft_step), - dtype=torch.float32, - device=target_next_token_ids0.device, - ) - if proposer.enable_dynmaic_mtp - else None - ) - if draft_step == 0: - return EagleSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=proposal_schedule_scores, - ) - - assert target_next_token_ids0.shape == (target_model_input0.batch_size,) - assert target_next_token_ids1.shape == (target_model_input1.batch_size,) - (accept_len0, accept_len1, req_start_rows0, req_start_rows1,) = prepare_dp_eagle_overlap_decode_inputs( - target_model_input0=target_model_input0, - target_model_input1=target_model_input1, - real_request_num0=real_request_num0, - real_request_num1=real_request_num1, - accept_len=accept_len, - ) - - model_inputs = [target_model_input0, target_model_input1] - model_outputs = [target_model_output0, target_model_output1] - target_token_ids_by_batch = [target_next_token_ids0, target_next_token_ids1] - req_start_rows_by_batch = [req_start_rows0, req_start_rows1] - accept_lens_by_batch = [accept_len0, accept_len1] - draft_inputs = [] - draft_token_ids_by_batch = [] - draft_hiddens_by_batch = [] - for (model_input, model_output, token_ids, req_start_rows, batch_accept_len, request_num,) in zip( - model_inputs, - model_outputs, - target_token_ids_by_batch, - req_start_rows_by_batch, - accept_lens_by_batch, - real_request_counts, - ): - selected_rows = select_accepted_tail_rows( - b_req_mtp_start_loc=req_start_rows, - accept_len=batch_accept_len, - input_ids=token_ids, - hidden=model_output.mtp_collector.spec_hidden, - b_req_idx=model_input.b_req_idx, - b_mtp_index=model_input.b_mtp_index, - b_seq_len=model_input.b_seq_len, - mem_indexes=model_input.mem_indexes, - b_shared_seq_len=model_input.b_shared_seq_len, - b_shared_radix_node_id=model_input.b_shared_radix_node_id, - b_position_delta=model_input.b_position_delta, - ) - draft_input = copy.copy(model_input) - draft_input.batch_size = request_num - draft_input.input_ids = selected_rows.input_ids - draft_input.b_req_idx = selected_rows.b_req_idx - draft_input.b_mtp_index = selected_rows.b_mtp_index - draft_input.b_seq_len = selected_rows.b_seq_len - draft_input.mem_indexes = selected_rows.mem_indexes - draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len - draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id - draft_input.b_position_delta = selected_rows.b_position_delta - draft_input.mem_indexes_cpu = None - draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(request_num)] - draft_inputs.append(draft_input) - draft_token_ids_by_batch.append(selected_rows.input_ids) - draft_hiddens_by_batch.append(selected_rows.hidden) - - proposal_row_offsets = (0, real_request_counts[0]) - for step in range(draft_step): - for batch_index, draft_input in enumerate(draft_inputs): - draft_input.input_ids = draft_token_ids_by_batch[batch_index] - draft_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] - - draft_outputs = get_draft_model(step)._microbatch_overlap_decode_cuda(*draft_inputs) - for batch_index, draft_output in enumerate(draft_outputs): - draft_token_ids, draft_token_probs = generate_eagle_tokens( - proposer=proposer, - model_output=draft_output, - map_draft_token_ids=map_draft_token_ids, - ) - draft_token_ids_by_batch[batch_index] = draft_token_ids - draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden - request_num = real_request_counts[batch_index] - proposal_row_start = proposal_row_offsets[batch_index] - proposal_token_ids[proposal_row_start : proposal_row_start + request_num, step] = draft_token_ids - if proposal_schedule_scores is not None: - proposal_schedule_scores[ - proposal_row_start : proposal_row_start + request_num, - step, - ] = draft_token_probs - - return EagleSpecProposal( - token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[], - schedule_scores=proposal_schedule_scores, - ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index 3a5c45f981..d582fe7143 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -1,16 +1,19 @@ +import copy + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids +from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( BaseDpOverlapProposer, ) -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( - fill_dp_eagle_draft_model_kv_state_overlap, - propose_next_dp_eagle_autoregressive_overlap, -) -from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( - EagleSpecProposal, +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.utils import ( + get_dp_overlap_req_start_rows, ) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import MtpMemIndexesToFree +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager class DpOverlapEagleWithAttProposer(BaseDpOverlapProposer): @@ -25,15 +28,29 @@ def fill_draft_model_kv_state_overlap( target_model_output1: ModelOutput, target_next_token_ids1: torch.Tensor, ) -> None: - fill_dp_eagle_draft_model_kv_state_overlap( - self, - target_model_input0, - target_model_output0, - target_next_token_ids0, - target_model_input1, - target_model_output1, - target_next_token_ids1, - ) + target_model_inputs = (target_model_input0, target_model_input1) + target_model_outputs = (target_model_output0, target_model_output1) + target_next_token_ids = (target_next_token_ids0, target_next_token_ids1) + assert len(self.backend.draft_models) == 1 + + draft_inputs = [] + for model_input, model_output, next_token_ids in zip( + target_model_inputs, + target_model_outputs, + target_next_token_ids, + ): + assert model_input.is_prefill + assert model_input.b_position_delta is None + assert next_token_ids.shape == model_input.b_req_idx.shape + draft_input = copy.copy(model_input) + self._prepare_eagle_prefill_inputs( + model_input=draft_input, + b_next_token_ids=next_token_ids, + mtp_draft_input_hiddens=model_output.mtp_collector.spec_hidden, + ) + draft_inputs.append(draft_input) + + self.backend.draft_models[0]._microbatch_overlap_prefill_cuda(*draft_inputs) def propose_next_overlap( self, @@ -48,17 +65,199 @@ def propose_next_overlap( accept_len: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: - return propose_next_dp_eagle_autoregressive_overlap( - proposer=self, - target_model_input0=target_model_input0, - target_model_output0=target_model_output0, - target_next_token_ids0=target_next_token_ids0, - target_model_input1=target_model_input1, - target_model_output1=target_model_output1, - target_next_token_ids1=target_next_token_ids1, - real_request_num0=real_request_num0, - real_request_num1=real_request_num1, - accept_len=accept_len, - draft_step=draft_step, - map_draft_token_ids=lambda token_ids: token_ids, + """提交两个 target verify microbatch 的 draft KV,并生成下一轮 proposal。""" + + assert draft_step > 0, "EAGLE attention requires draft_step to be greater than 0" + assert not target_model_input0.is_prefill + assert not target_model_input1.is_prefill + assert len(self.backend.draft_models) == 1 + + req_num_by_batch = (int(real_request_num0), int(real_request_num1)) + req_num = sum(req_num_by_batch) + assert accept_len.shape == (req_num,) + assert target_next_token_ids0.shape == (target_model_input0.batch_size,) + assert target_next_token_ids1.shape == (target_model_input1.batch_size,) + assert target_model_output0.mtp_collector.spec_hidden.shape[0] == target_model_input0.batch_size + assert target_model_output1.mtp_collector.spec_hidden.shape[0] == target_model_input1.batch_size + assert target_model_input0.b_position_delta is not None + assert target_model_input1.b_position_delta is not None + + accept_len_by_batch = ( + accept_len[: req_num_by_batch[0]], + accept_len[req_num_by_batch[0] :], + ) + b_req_mtp_start_loc = ( + get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input0.b_mtp_index, + req_num=req_num_by_batch[0], + ), + get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input1.b_mtp_index, + req_num=req_num_by_batch[1], + ), + ) + accepted_tail_rows_by_batch = tuple( + (req_mtp_start_loc + batch_accept_len - 1).long() + for req_mtp_start_loc, batch_accept_len in zip( + b_req_mtp_start_loc, + accept_len_by_batch, + ) + ) + + model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] + model_outputs = (target_model_output0, target_model_output1) + target_next_token_ids_by_batch = (target_next_token_ids0, target_next_token_ids1) + position_deltas_by_batch = tuple(model_input.b_position_delta for model_input in model_inputs) + max_kv_seq_lens_by_batch = tuple(model_input.max_kv_seq_len for model_input in model_inputs) + for model_input, model_output, token_ids in zip( + model_inputs, + model_outputs, + target_next_token_ids_by_batch, + ): + model_input.input_ids = token_ids + model_input.mtp_draft_input_hiddens = model_output.mtp_collector.spec_hidden + + proposal_token_ids = target_next_token_ids0.new_empty((req_num, draft_step)) + schedule_scores = ( + torch.empty( + (req_num, draft_step), + dtype=torch.float32, + device=target_next_token_ids0.device, + ) + if self.enable_dynmaic_mtp + else None + ) + + draft_model = self.backend.draft_models[0] + extend_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) + draft_token_ids_by_batch = [] + draft_hiddens_by_batch = [] + draft_seq_lens_by_batch = [] + draft_req_indices_by_batch = [] + draft_shared_seq_lens_by_batch = [] + draft_shared_radix_node_ids_by_batch = [] + proposal_row_offsets = (0, req_num_by_batch[0]) + + for batch_index, (model_input, extend_output, accepted_tail_rows, batch_req_num) in enumerate( + zip( + model_inputs, + extend_outputs, + accepted_tail_rows_by_batch, + req_num_by_batch, + ) + ): + accepted_tail_output = ModelOutput(logits=extend_output.logits.index_select(0, accepted_tail_rows)) + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(accepted_tail_output) + draft_token_probs = draft_token_probs.float() + else: + draft_token_ids = self._gen_argmax_token_ids(accepted_tail_output) + + draft_token_ids_by_batch.append(draft_token_ids) + draft_hiddens_by_batch.append(extend_output.mtp_collector.spec_hidden.index_select(0, accepted_tail_rows)) + draft_seq_lens_by_batch.append(model_input.b_seq_len.index_select(0, accepted_tail_rows) + 1) + draft_req_indices_by_batch.append(model_input.b_req_idx.index_select(0, accepted_tail_rows)) + draft_shared_seq_lens_by_batch.append(model_input.b_shared_seq_len.index_select(0, accepted_tail_rows)) + draft_shared_radix_node_ids_by_batch.append( + model_input.b_shared_radix_node_id.index_select(0, accepted_tail_rows) + ) + + proposal_row_start = proposal_row_offsets[batch_index] + proposal_row_end = proposal_row_start + batch_req_num + proposal_token_ids[proposal_row_start:proposal_row_end, 0] = draft_token_ids + if schedule_scores is not None: + schedule_scores[proposal_row_start:proposal_row_end, 0] = draft_token_probs + + if draft_step == 1: + return EagleSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, + ) + + for batch_index, model_input in enumerate(model_inputs): + model_input.is_prefill = False + model_input.batch_size = req_num_by_batch[batch_index] + model_input.b_req_idx = draft_req_indices_by_batch[batch_index] + model_input.b_mtp_index = torch.zeros_like(model_input.b_req_idx) + model_input.b_seq_len = draft_seq_lens_by_batch[batch_index] + model_input.b_position_delta = position_deltas_by_batch[batch_index].index_select( + 0, accepted_tail_rows_by_batch[batch_index] + ) + model_input.b_shared_seq_len = draft_shared_seq_lens_by_batch[batch_index] + model_input.b_shared_radix_node_id = draft_shared_radix_node_ids_by_batch[batch_index] + if len(model_input.multimodal_params) != model_input.batch_size: + empty_multimodal_params = {"images": [], "audios": []} + model_input.multimodal_params = [empty_multimodal_params] * model_input.batch_size + + extra_mem_indexes_cpu = mtp_utils.alloc_mem_indexes(req_num * (draft_step - 1)) + extra_mem_indexes = extra_mem_indexes_cpu.to(device=target_next_token_ids0.device, non_blocking=True) + + for step in range(1, draft_step): + mem_start = (step - 1) * req_num + step_mem_indexes = extra_mem_indexes[mem_start : mem_start + req_num] + mem_offset = 0 + for batch_index, model_input in enumerate(model_inputs): + batch_req_num = req_num_by_batch[batch_index] + model_input.input_ids = draft_token_ids_by_batch[batch_index] + model_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] + model_input.mem_indexes = step_mem_indexes[mem_offset : mem_offset + batch_req_num] + model_input.max_kv_seq_len = max_kv_seq_lens_by_batch[batch_index] + step + model_input.total_token_num = model_input.batch_size * model_input.max_kv_seq_len + mem_offset += batch_req_num + + draft_outputs = draft_model._microbatch_overlap_decode_cuda(*model_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self._gen_argmax_token_ids_and_prob(draft_output) + draft_token_probs = draft_token_probs.float() + else: + draft_token_ids = self._gen_argmax_token_ids(draft_output) + draft_token_ids_by_batch[batch_index] = draft_token_ids + draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden + draft_seq_lens_by_batch[batch_index].add_(1) + + batch_req_num = req_num_by_batch[batch_index] + proposal_row_start = proposal_row_offsets[batch_index] + proposal_row_end = proposal_row_start + batch_req_num + proposal_token_ids[proposal_row_start:proposal_row_end, step] = draft_token_ids + if schedule_scores is not None: + schedule_scores[proposal_row_start:proposal_row_end, step] = draft_token_probs + + return EagleSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], + schedule_scores=schedule_scores, + ) + + def _gen_argmax_token_ids(self, model_output: ModelOutput) -> torch.Tensor: + """生成 target vocabulary 下的候选 token;EAGLE3 会覆盖词表映射。""" + + return self.backend._gen_argmax_token_ids(model_output) + + def _gen_argmax_token_ids_and_prob(self, model_output: ModelOutput): + """生成候选 token 及其概率;EAGLE3 会覆盖 token 的词表映射。""" + + return self.backend._gen_argmax_token_ids_and_prob(model_output) + + @staticmethod + def _prepare_eagle_prefill_inputs( + model_input: ModelInput, + b_next_token_ids: torch.Tensor, + mtp_draft_input_hiddens: torch.Tensor, + ) -> None: + """构造 DP-overlap EAGLE draft model 的 prefill 输入。""" + + model_input.b_is_decode_req = g_pin_mem_manager.get_const_gpu_tensor( + key="dp_overlap_eagle_prefill_b_is_decode_req", + shape=model_input.b_req_idx.shape, + fill_value=False, + dtype=torch.bool, + ) + model_input.input_ids = gen_mtp_new_input_ids( + input_ids=model_input.input_ids, + b_next_token_ids=b_next_token_ids, + b_seq_len=model_input.b_seq_len, + b_ready_cache_len=model_input.b_ready_cache_len, ) + model_input.mtp_draft_input_hiddens = mtp_draft_input_hiddens diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/utils.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/utils.py new file mode 100644 index 0000000000..96393c0dcc --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/utils.py @@ -0,0 +1,17 @@ +import torch + +from lightllm.common.basemodel.triton_kernel.mtp_utils import gen_b_req_mtp_start_loc + + +def get_dp_overlap_req_start_rows(b_mtp_index: torch.Tensor, req_num: int) -> torch.Tensor: + """根据动态 verify 布局恢复请求起始行。""" + + req_num = int(req_num) + if req_num == 0: + assert b_mtp_index.numel() == 0 + return torch.empty((0,), dtype=torch.int32, device=b_mtp_index.device) + assert b_mtp_index.is_cuda, "b_mtp_index must be a CUDA tensor" + return gen_b_req_mtp_start_loc( + b_mtp_index=b_mtp_index, + num_reqs=req_num, + ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index 5c0607c082..ca003c7ab4 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -1,11 +1,14 @@ +import copy + import torch from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.triton_kernel.select_mtp_rows import select_accepted_tail_rows from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( BaseDpOverlapProposer, ) -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( - propose_next_dp_no_att_overlap, +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.utils import ( + get_dp_overlap_req_start_rows, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( VanillaSpecProposal, @@ -39,23 +42,113 @@ def propose_next_overlap( accept_len: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: - proposal = propose_next_dp_no_att_overlap( - proposer=self, - target_model_input0=target_model_input0, - target_model_output0=target_model_output0, - target_next_token_ids0=target_next_token_ids0, - target_model_input1=target_model_input1, - target_model_output1=target_model_output1, - target_next_token_ids1=target_next_token_ids1, - real_request_num0=real_request_num0, - real_request_num1=real_request_num1, - accept_len=accept_len, - draft_step=draft_step, - get_draft_model=lambda step: self.backend.draft_models[step], - map_draft_token_ids=lambda token_ids: token_ids, + req_num_by_batch = (int(real_request_num0), int(real_request_num1)) + req_num = sum(req_num_by_batch) + assert accept_len.shape == (req_num,) + + proposal_token_ids = target_next_token_ids0.new_empty((req_num, draft_step)) + schedule_scores = ( + torch.empty( + (req_num, draft_step), + dtype=torch.float32, + device=target_next_token_ids0.device, + ) + if self.enable_dynmaic_mtp + else None + ) + + if draft_step == 0: + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, + ) + + assert target_next_token_ids0.shape == (target_model_input0.batch_size,) + assert target_next_token_ids1.shape == (target_model_input1.batch_size,) + accept_len_by_batch = ( + accept_len[: req_num_by_batch[0]], + accept_len[req_num_by_batch[0] :], + ) + b_req_mtp_start_loc = ( + get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input0.b_mtp_index, + req_num=req_num_by_batch[0], + ), + get_dp_overlap_req_start_rows( + b_mtp_index=target_model_input1.b_mtp_index, + req_num=req_num_by_batch[1], + ), ) + + model_inputs = (target_model_input0, target_model_input1) + model_outputs = (target_model_output0, target_model_output1) + target_next_token_ids_by_batch = (target_next_token_ids0, target_next_token_ids1) + draft_inputs = [] + draft_token_ids_by_batch = [] + draft_hiddens_by_batch = [] + for model_input, model_output, token_ids, req_mtp_start_loc, batch_accept_len, batch_req_num in zip( + model_inputs, + model_outputs, + target_next_token_ids_by_batch, + b_req_mtp_start_loc, + accept_len_by_batch, + req_num_by_batch, + ): + selected_rows = select_accepted_tail_rows( + b_req_mtp_start_loc=req_mtp_start_loc, + accept_len=batch_accept_len, + input_ids=token_ids, + hidden=model_output.mtp_collector.spec_hidden, + b_req_idx=model_input.b_req_idx, + b_mtp_index=model_input.b_mtp_index, + b_seq_len=model_input.b_seq_len, + mem_indexes=model_input.mem_indexes, + b_shared_seq_len=model_input.b_shared_seq_len, + b_shared_radix_node_id=model_input.b_shared_radix_node_id, + b_position_delta=model_input.b_position_delta, + ) + draft_input = copy.copy(model_input) + draft_input.batch_size = batch_req_num + draft_input.input_ids = selected_rows.input_ids + draft_input.b_req_idx = selected_rows.b_req_idx + draft_input.b_mtp_index = selected_rows.b_mtp_index + draft_input.b_seq_len = selected_rows.b_seq_len + draft_input.mem_indexes = selected_rows.mem_indexes + draft_input.b_shared_seq_len = selected_rows.b_shared_seq_len + draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id + draft_input.b_position_delta = selected_rows.b_position_delta + draft_input.mem_indexes_cpu = None + draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(batch_req_num)] + draft_inputs.append(draft_input) + draft_token_ids_by_batch.append(selected_rows.input_ids) + draft_hiddens_by_batch.append(selected_rows.hidden) + + proposal_row_offsets = (0, req_num_by_batch[0]) + for step in range(draft_step): + for batch_index, draft_input in enumerate(draft_inputs): + draft_input.input_ids = draft_token_ids_by_batch[batch_index] + draft_input.mtp_draft_input_hiddens = draft_hiddens_by_batch[batch_index] + + draft_outputs = self.backend.draft_models[step]._microbatch_overlap_decode_cuda(*draft_inputs) + for batch_index, draft_output in enumerate(draft_outputs): + if self.enable_dynmaic_mtp: + draft_token_ids, draft_token_probs = self.backend._gen_argmax_token_ids_and_prob(draft_output) + draft_token_probs = draft_token_probs.float() + else: + draft_token_ids = self.backend._gen_argmax_token_ids(draft_output) + draft_token_ids_by_batch[batch_index] = draft_token_ids + draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden + + batch_req_num = req_num_by_batch[batch_index] + proposal_row_start = proposal_row_offsets[batch_index] + proposal_row_end = proposal_row_start + batch_req_num + proposal_token_ids[proposal_row_start:proposal_row_end, step] = draft_token_ids + if schedule_scores is not None: + schedule_scores[proposal_row_start:proposal_row_end, step] = draft_token_probs + return VanillaSpecProposal( - token_ids=proposal.token_ids, - extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, - schedule_scores=proposal.schedule_scores, + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=schedule_scores, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index 1a6b8dec4d..ee4fa89abb 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -12,7 +12,7 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( BaseDpOverlapProposer, ) -from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_utils import ( +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.utils import ( get_dp_overlap_req_start_rows, ) from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( @@ -80,21 +80,27 @@ def propose_next_overlap( assert target_next_token_ids0.shape == (target_model_input0.batch_size,) assert target_next_token_ids1.shape == (target_model_input1.batch_size,) - req_start_rows = ( + b_req_mtp_start_loc = ( get_dp_overlap_req_start_rows( b_mtp_index=target_model_input0.b_mtp_index, - request_num=real_request_num0, + req_num=real_request_num0, ), get_dp_overlap_req_start_rows( b_mtp_index=target_model_input1.b_mtp_index, - request_num=real_request_num1, + req_num=real_request_num1, ), ) accept_len_by_batch = ( accept_len[: req_num_by_batch[0]], accept_len[req_num_by_batch[0] : sum(req_num_by_batch)], ) - accepted_tail_rows = tuple(starts + lengths - 1 for starts, lengths in zip(req_start_rows, accept_len_by_batch)) + accepted_tail_rows = tuple( + req_mtp_start_loc + batch_accept_len - 1 + for req_mtp_start_loc, batch_accept_len in zip( + b_req_mtp_start_loc, + accept_len_by_batch, + ) + ) model_inputs = [copy.copy(target_model_input0), copy.copy(target_model_input1)] draft_token_ids = [target_next_token_ids0, target_next_token_ids1] draft_hiddens = [ @@ -145,7 +151,7 @@ def propose_next_overlap( draft_token_ids[batch_index] = overlay_chained_mtp_decode_input( input_ids=model_input.input_ids, draft_token_ids=draft_token_ids[batch_index], - b_req_mtp_start_loc=req_start_rows[batch_index], + b_req_mtp_start_loc=b_req_mtp_start_loc[batch_index], accept_len=accept_len_by_batch[batch_index], ) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 7a4ee8d889..6e35a37185 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -5,6 +5,10 @@ from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers import eagle_with_att +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.utils import ( + get_dp_overlap_req_start_rows, +) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle3 import ( DpOverlapEagle3Proposer, ) @@ -83,7 +87,25 @@ def _target_input(batch_size, b_mtp_index=None, device="cpu"): ) +def _patch_cpu_req_start_rows(monkeypatch): + def get_cpu_req_start_rows(b_mtp_index, req_num): + req_start_rows = torch.nonzero(b_mtp_index == 0, as_tuple=False).flatten().to(dtype=torch.int32) + assert req_start_rows.shape == (req_num,) + return req_start_rows + + monkeypatch.setattr(eagle_with_att, "get_dp_overlap_req_start_rows", get_cpu_req_start_rows) + + +def test_dp_overlap_req_start_rows_rejects_nonempty_cpu_input(): + with pytest.raises(AssertionError, match="must be a CUDA tensor"): + get_dp_overlap_req_start_rows( + b_mtp_index=torch.tensor([0, 1], dtype=torch.int32), + req_num=1, + ) + + def test_overlap_eagle_supports_variable_verify_layout(monkeypatch): + _patch_cpu_req_start_rows(monkeypatch) draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -180,6 +202,7 @@ def test_overlap_eagle_supports_empty_verify_rows(monkeypatch): def test_overlap_eagle_returns_dynamic_schedule_scores(monkeypatch): + _patch_cpu_req_start_rows(monkeypatch) draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -288,6 +311,7 @@ def test_overlap_eagle_no_att_supports_dynamic_draft_step(): def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): + _patch_cpu_req_start_rows(monkeypatch) draft_model = _DraftModel() backend = SimpleNamespace( max_draft_step=2, @@ -360,3 +384,22 @@ def test_eagle3_maps_draft_token_ids_in_proposer(): assert torch.equal(token_ids, torch.tensor([101, 102])) assert torch.equal(token_ids_with_prob, torch.tensor([103, 104])) assert torch.equal(probs, torch.tensor([0.8, 0.7])) + + +def test_dp_overlap_eagle3_maps_draft_token_ids_in_proposer(): + proposer = DpOverlapEagle3Proposer.__new__(DpOverlapEagle3Proposer) + proposer.backend = SimpleNamespace( + draft_models=[SimpleNamespace(map_draft_vocab_to_main_vocab=lambda token_ids: token_ids + 100)], + _gen_argmax_token_ids=lambda _: torch.tensor([1, 2]), + _gen_argmax_token_ids_and_prob=lambda _: ( + torch.tensor([3, 4]), + torch.tensor([0.8, 0.7]), + ), + ) + + token_ids = proposer._gen_argmax_token_ids(ModelOutput(logits=torch.empty(0))) + token_ids_with_prob, probs = proposer._gen_argmax_token_ids_and_prob(ModelOutput(logits=torch.empty(0))) + + assert torch.equal(token_ids, torch.tensor([101, 102])) + assert torch.equal(token_ids_with_prob, torch.tensor([103, 104])) + assert torch.equal(probs, torch.tensor([0.8, 0.7])) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py index 54c7a0006f..7baf34061b 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -371,7 +371,6 @@ def test_each_mode_proposer_inherits_its_expected_implementation_base(): DpOverlapVanillaNoAttProposer, DpOverlapEagleWithAttProposer, DpOverlapEagleNoAttProposer, - DpOverlapEagle3Proposer, ) for proposer_type in proposer_types: @@ -379,6 +378,7 @@ def test_each_mode_proposer_inherits_its_expected_implementation_base(): assert Eagle3Proposer.__bases__ == (EagleWithAttProposer,) for proposer_type in dp_overlap_proposer_types: assert proposer_type.__bases__ == (BaseDpOverlapProposer,) + assert DpOverlapEagle3Proposer.__bases__ == (DpOverlapEagleWithAttProposer,) def test_each_mtp_mode_builds_its_own_proposer(): From 414c39b19b6284793af5b44eaa6209e339278ed5 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sat, 22 Aug 2026 09:28:33 +0000 Subject: [PATCH 091/103] fix: honor qwen3 eagle attention head dim --- .../layer_infer/transformer_layer_infer.py | 4 ++-- .../models/test_qwen3_eagle_model_input.py | 19 +++++++++++++++++++ 2 files changed, 21 insertions(+), 2 deletions(-) diff --git a/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py index f362c6af8a..4f34a9e76b 100644 --- a/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py @@ -1,12 +1,12 @@ import torch from lightllm.common.basemodel.infer_struct import InferStateInfo -from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd +from lightllm.models.qwen3.layer_infer.transformer_layer_infer import Qwen3TransformerLayerInfer from lightllm.models.qwen3_eagle.layer_weights.transformer_layer_weight import Qwen3EagleTransformerLayerWeight -class Qwen3EagleTransformerLayerInfer(LlamaTransformerLayerInfer): +class Qwen3EagleTransformerLayerInfer(Qwen3TransformerLayerInfer): def _get_qkv(self, input, infer_state: InferStateInfo, layer_weight: Qwen3EagleTransformerLayerWeight): input_part = self._att_norm(input, infer_state, layer_weight) target_part = layer_weight.hidden_norm_weight_( diff --git a/unit_tests/models/test_qwen3_eagle_model_input.py b/unit_tests/models/test_qwen3_eagle_model_input.py index 5ffd9f73fb..1395153351 100644 --- a/unit_tests/models/test_qwen3_eagle_model_input.py +++ b/unit_tests/models/test_qwen3_eagle_model_input.py @@ -3,11 +3,30 @@ import pytest import torch +from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.llama.model import LlamaTpPartModel from lightllm.models.qwen3_eagle.layer_infer.pre_layer_infer import Qwen3EaglePreLayerInfer +from lightllm.models.qwen3_eagle.layer_infer.transformer_layer_infer import Qwen3EagleTransformerLayerInfer from lightllm.models.qwen3_eagle.model import Qwen3EagleModel +def test_qwen3_eagle_uses_configured_attention_head_dim(monkeypatch): + def init_llama_layer(self, layer_num, network_config): + self.head_dim_ = network_config["hidden_size"] // network_config["num_attention_heads"] + + monkeypatch.setattr(LlamaTransformerLayerInfer, "__init__", init_llama_layer) + layer = Qwen3EagleTransformerLayerInfer( + layer_num=0, + network_config={ + "hidden_size": 2560, + "num_attention_heads": 32, + "head_dim": 128, + }, + ) + + assert layer.head_dim_ == 128 + + def test_qwen3_eagle_projects_concatenated_target_hiddens_before_inferstate(monkeypatch): parent_calls = [] monkeypatch.setattr( From e56d533a1a73400a21bb5ab8bdbe2f7b8d36f60c Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 09:49:33 +0000 Subject: [PATCH 092/103] Use workspace buffers for FA3 page tables --- lightllm/common/basemodel/attention/fa3/fp.py | 10 +++++----- lightllm/common/basemodel/attention/fa3/mla.py | 10 +++++----- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/lightllm/common/basemodel/attention/fa3/fp.py b/lightllm/common/basemodel/attention/fa3/fp.py index ea11b8307d..0f0f62ec63 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -2,7 +2,6 @@ import torch from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl from typing import Optional, TYPE_CHECKING -from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.sgl_utils import flash_attn_with_kvcache, flash_attn_with_kvcache_autotune from lightllm.common.basemodel.triton_kernel.fa3_utils import ( build_dynamic_spec_fa3_decode_params, @@ -30,13 +29,14 @@ def get_page_table_buffer(self): max_att_batch_size //= model.mtp_manager.get_decode_batch_multiplier(model.is_mtp_draft_model) buffer_count = 2 if model.args.enable_decode_microbatch_overlap else 1 + workspace_size = max_att_batch_size * model.graph_max_len_in_batch self._shared_page_table_buffer = [ - torch.empty( - max_att_batch_size * model.graph_max_len_in_batch, + self.get_gpu_workspace_buffer( + key_name=f"fa3_fp_page_table_{buffer_index}", + workspace_size=workspace_size, dtype=torch.int32, - device=get_current_device_id(), ) - for _ in range(buffer_count) + for buffer_index in range(buffer_count) ] return self._shared_page_table_buffer diff --git a/lightllm/common/basemodel/attention/fa3/mla.py b/lightllm/common/basemodel/attention/fa3/mla.py index 3f4a9170a6..5b757a0d67 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -2,7 +2,6 @@ import torch from ..base_att import BaseAttBackend, BasePrefillAttState, BaseDecodeAttState, AttControl from typing import Optional, TYPE_CHECKING, Tuple -from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.sgl_utils import flash_attn_with_kvcache from lightllm.common.basemodel.triton_kernel.fa3_utils import build_dynamic_spec_fa3_decode_params, page_table_copy from lightllm.common.basemodel.triton_kernel.gen_prefill_params import gen_cumsum_pad0_tensor @@ -28,13 +27,14 @@ def get_page_table_buffer(self): max_att_batch_size //= model.mtp_manager.get_decode_batch_multiplier(model.is_mtp_draft_model) buffer_count = 2 if model.args.enable_decode_microbatch_overlap else 1 + workspace_size = max_att_batch_size * model.graph_max_len_in_batch self._shared_page_table_buffer = [ - torch.empty( - max_att_batch_size * model.graph_max_len_in_batch, + self.get_gpu_workspace_buffer( + key_name=f"fa3_mla_page_table_{buffer_index}", + workspace_size=workspace_size, dtype=torch.int32, - device=get_current_device_id(), ) - for _ in range(buffer_count) + for buffer_index in range(buffer_count) ] return self._shared_page_table_buffer From e745657900535e8663f0e60fc22a460917b42242 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 12:05:38 +0000 Subject: [PATCH 093/103] Rename chained MTP decode input builder --- ...t.py => build_chained_mtp_decode_input.py} | 6 +- .../dp_overlap_proposers/vanilla_with_att.py | 6 +- .../proposers/vanilla_with_att.py | 6 +- .../test_build_chained_mtp_decode_input.py | 94 +++++++++++++++++++ 4 files changed, 104 insertions(+), 8 deletions(-) rename lightllm/common/basemodel/triton_kernel/{overlay_mtp_decode_input.py => build_chained_mtp_decode_input.py} (94%) create mode 100644 test/kernel/test_build_chained_mtp_decode_input.py diff --git a/lightllm/common/basemodel/triton_kernel/overlay_mtp_decode_input.py b/lightllm/common/basemodel/triton_kernel/build_chained_mtp_decode_input.py similarity index 94% rename from lightllm/common/basemodel/triton_kernel/overlay_mtp_decode_input.py rename to lightllm/common/basemodel/triton_kernel/build_chained_mtp_decode_input.py index 78a57c004f..80da12fa5c 100644 --- a/lightllm/common/basemodel/triton_kernel/overlay_mtp_decode_input.py +++ b/lightllm/common/basemodel/triton_kernel/build_chained_mtp_decode_input.py @@ -4,7 +4,7 @@ @triton.jit -def _overlay_chained_mtp_decode_input_kernel( +def _build_chained_mtp_decode_input_kernel( input_ids, draft_token_ids, b_req_mtp_start_loc, @@ -24,7 +24,7 @@ def _overlay_chained_mtp_decode_input_kernel( @torch.no_grad() -def overlay_chained_mtp_decode_input( +def build_chained_mtp_decode_input_inplace( input_ids: torch.Tensor, draft_token_ids: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, @@ -62,7 +62,7 @@ def overlay_chained_mtp_decode_input( if req_num == 0: return draft_token_ids - _overlay_chained_mtp_decode_input_kernel[(req_num,)]( + _build_chained_mtp_decode_input_kernel[(req_num,)]( input_ids=input_ids, draft_token_ids=draft_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index ee4fa89abb..ecda2dfb21 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -6,8 +6,8 @@ from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import ( gen_mtp_new_input_ids, ) -from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import ( - overlay_chained_mtp_decode_input, +from lightllm.common.basemodel.triton_kernel.build_chained_mtp_decode_input import ( + build_chained_mtp_decode_input_inplace, ) from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( BaseDpOverlapProposer, @@ -148,7 +148,7 @@ def propose_next_overlap( if step + 1 < draft_step: for batch_index, model_input in enumerate(model_inputs): - draft_token_ids[batch_index] = overlay_chained_mtp_decode_input( + draft_token_ids[batch_index] = build_chained_mtp_decode_input_inplace( input_ids=model_input.input_ids, draft_token_ids=draft_token_ids[batch_index], b_req_mtp_start_loc=b_req_mtp_start_loc[batch_index], diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py index 87a8376848..ac844cc8d6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -4,7 +4,9 @@ from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids -from lightllm.common.basemodel.triton_kernel.overlay_mtp_decode_input import overlay_chained_mtp_decode_input +from lightllm.common.basemodel.triton_kernel.build_chained_mtp_decode_input import ( + build_chained_mtp_decode_input_inplace, +) from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager @@ -103,7 +105,7 @@ def propose_next( # 下一层不能直接使用所有行的 draft 预测。已接受前缀继续使用 # main/上一层输入中的真实 token,仅在 tail 行接上本级新生成 # 的 draft token,从而形成逐级左移并覆盖尾部的级联输入。 - draft_token_ids = overlay_chained_mtp_decode_input( + draft_token_ids = build_chained_mtp_decode_input_inplace( input_ids=draft_input.input_ids, draft_token_ids=draft_token_ids, b_req_mtp_start_loc=b_req_mtp_start_loc, diff --git a/test/kernel/test_build_chained_mtp_decode_input.py b/test/kernel/test_build_chained_mtp_decode_input.py new file mode 100644 index 0000000000..90c9465251 --- /dev/null +++ b/test/kernel/test_build_chained_mtp_decode_input.py @@ -0,0 +1,94 @@ +import pytest +import torch + +from lightllm.common.basemodel.triton_kernel.build_chained_mtp_decode_input import ( + build_chained_mtp_decode_input_inplace, +) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_build_chained_mtp_decode_input_inplace_shifts_accepted_tokens_and_keeps_tail(): + input_ids = torch.tensor( + [ + 10, + 11, + 12, + 20, + 30, + 31, + 40, + 41, + 42, + 43, + 90, + 91, + ], + dtype=torch.int64, + device="cuda", + ) + draft_token_ids = torch.tensor( + [ + 100, + 101, + 102, + 200, + 300, + 301, + 400, + 401, + 402, + 403, + 900, + 901, + ], + dtype=torch.int64, + device="cuda", + ) + b_req_mtp_start_loc = torch.tensor([0, 3, 4, 6], dtype=torch.int32, device="cuda") + accept_len = torch.tensor([3, 1, 2, 4], dtype=torch.int32, device="cuda") + + result = build_chained_mtp_decode_input_inplace( + input_ids=input_ids, + draft_token_ids=draft_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + ) + + expected = torch.tensor( + [ + 11, + 12, + 102, + 200, + 31, + 301, + 41, + 42, + 43, + 403, + 900, + 901, + ], + dtype=torch.int64, + device="cuda", + ) + assert result is draft_token_ids + torch.testing.assert_close(draft_token_ids, expected) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_build_chained_mtp_decode_input_inplace_handles_empty_request_set(): + input_ids = torch.empty((0,), dtype=torch.int64, device="cuda") + draft_token_ids = torch.empty((0,), dtype=torch.int64, device="cuda") + b_req_mtp_start_loc = torch.empty((0,), dtype=torch.int32, device="cuda") + accept_len = torch.empty((0,), dtype=torch.int32, device="cuda") + + result = build_chained_mtp_decode_input_inplace( + input_ids=input_ids, + draft_token_ids=draft_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + accept_len=accept_len, + ) + + assert result is draft_token_ids + assert draft_token_ids.numel() == 0 From 64515896424a7da4c432f1e3e26ea2c4fbd4154e Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 14:17:44 +0000 Subject: [PATCH 094/103] Fix MTP sampling buffer width --- lightllm/common/basemodel/basemodel.py | 7 +++---- lightllm/common/req_manager.py | 3 ++- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 6ffc3f3535..f2b6bae085 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -117,10 +117,9 @@ def __init__(self, kvargs): self._init_quant() enable_weight_cpu_backup = self.args.enable_weight_cpu_backup - with profile_mtp_weight_memory(self), self.torch_memory_saver.region( - tag=MemoryTag.WEIGHT, enable_cpu_backup=enable_weight_cpu_backup - ): - self._init_weights() + with profile_mtp_weight_memory(self): + with self.torch_memory_saver.region(tag=MemoryTag.WEIGHT, enable_cpu_backup=enable_weight_cpu_backup): + self._init_weights() with self.torch_memory_saver.region(tag=MemoryTag.KV_CACHE): self._init_req_manager() self._init_mem_manager() diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 7ff84fd058..3de7de8f12 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -113,11 +113,12 @@ def __init__(self, max_request_num): # mode ["cpu_counter", "pin_mem_counter", "gpu_counter"] self.penalty_counter_mode = get_env_start_args().penalty_counter_mode self.vocab_size = get_vocab_size(get_env_start_args().model_dir) + self.mtp_verify_width = get_env_start_args().mtp_step + 1 self.req_to_presence_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") self.req_to_frequency_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") self.req_to_repetition_penalty = torch.zeros(max_request_num + 1, dtype=torch.float32, device="cuda") self.req_to_next_token_ids = torch.zeros( - (max_request_num + 1, 8), + (max_request_num + 1, self.mtp_verify_width), dtype=torch.int64, device="cuda", ) From e7c80c8769310b0c4e9458f8921c50063a1af4dc Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 14:22:43 +0000 Subject: [PATCH 095/103] Restrict diverse MTP modes to no-attention --- .../router/model_infer/mode_backend/diverse_backend/impl.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py index 4802b83235..21979cbef0 100644 --- a/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/diverse_backend/impl.py @@ -25,8 +25,6 @@ def __init__(self) -> None: spec_mode = get_env_start_args().mtp_mode if spec_mode is not None: assert spec_mode in [ - "vanilla_with_att", - "eagle_with_att", "vanilla_no_att", "eagle_no_att", ] From 9bb1a274ab151d95c0711702876715228fafb8e4 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 14:35:13 +0000 Subject: [PATCH 096/103] Simplify DP overlap prefill token slicing --- .../mode_backend/dp_backend/impl.py | 36 ++----------------- .../test_dp_overlap_spec_engine.py | 11 ------ 2 files changed, 2 insertions(+), 45 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index de93987c43..80dee159d0 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -82,26 +82,6 @@ def init_spec_engine(self): ) return - @staticmethod - def _build_padded_next_token_ids( - token_ids: torch.Tensor | None, - batch_size: int, - copy_len: int, - device: torch.device, - source_start: int = 0, - ) -> torch.Tensor: - """Pad DP-local draft tokens to the collective batch shape.""" - - copy_len = int(copy_len) - source_start = int(source_start) - padded_token_ids = torch.zeros((int(batch_size),), dtype=torch.int64, device=device) - if copy_len > 0: - padded_token_ids[:copy_len].copy_( - token_ids[source_start : source_start + copy_len], - non_blocking=True, - ) - return padded_token_ids - def _init_reqs(self, reqs: List[Tuple]): if not self.args.enable_dp_prompt_cache_fetch: return super()._init_reqs(reqs) @@ -719,20 +699,8 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I else: next_token_ids = torch.empty((0,), dtype=torch.int64, device=logits.device) - target_next_token_ids_gpu0 = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input0.batch_size, - copy_len=req_num0, - device=model_input0.b_req_idx.device, - source_start=0, - ) - target_next_token_ids_gpu1 = self._build_padded_next_token_ids( - token_ids=next_token_ids, - batch_size=model_input1.batch_size, - copy_len=req_num1, - device=model_input1.b_req_idx.device, - source_start=req_num0, - ) + target_next_token_ids_gpu0 = next_token_ids[:req_num0] + target_next_token_ids_gpu1 = next_token_ids[req_num0:] self.prefill_draft_engine.fill_draft_model_kv_state_overlap( target_model_input0=model_input0, diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 2bc7e809e3..66543e68c9 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -230,17 +230,6 @@ def test_dp_prefill_and_decode_select_overlap_engine_independently(): assert backend.decode_draft_engine is expected_decode_engine -def test_padded_token_ids_support_empty_dp_rank(): - padded_token_ids = DPChunkedPrefillBackend._build_padded_next_token_ids( - token_ids=None, - batch_size=4, - copy_len=0, - device=torch.device("cpu"), - ) - - assert torch.equal(padded_token_ids, torch.zeros(4, dtype=torch.int64)) - - @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_dp_decode_mtp_runs_common_engine_for_empty_batch(monkeypatch): device = "cuda" From 969ba1aa7948749890424484e3a4557f641d2fea Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 14:43:47 +0000 Subject: [PATCH 097/103] Simplify DP prefill real-shape handling --- .../mode_backend/dp_backend/impl.py | 47 ++++++------------- 1 file changed, 15 insertions(+), 32 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 80dee159d0..09a2d96c92 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -292,15 +292,13 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer req_num0, req_num1 = len(run_reqs0), len(run_reqs1) logits = torch.empty((req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device) - logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) - logits[req_num0 : req_num0 + req_num1, :].copy_(logits1[0:req_num1, :], non_blocking=True) + logits[0:req_num0, :].copy_(logits0, non_blocking=True) + logits[req_num0 : req_num0 + req_num1, :].copy_(logits1, non_blocking=True) run_reqs = run_reqs0 + run_reqs1 - b_has_out_cpu = ( - model_input0.b_prefill_has_output_cpu[0:req_num0] + model_input1.b_prefill_has_output_cpu[0:req_num1] - ) - b_mtp_index = torch.cat((model_input0.b_mtp_index[0:req_num0], model_input1.b_mtp_index[0:req_num1]), dim=0) - b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) + b_has_out_cpu = model_input0.b_prefill_has_output_cpu + model_input1.b_prefill_has_output_cpu + b_mtp_index = torch.cat((model_input0.b_mtp_index, model_input1.b_mtp_index), dim=0) + b_req_idx = torch.cat((model_input0.b_req_idx, model_input1.b_req_idx), dim=0) if req_num0 + req_num1 > 0: ( @@ -411,10 +409,10 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] req_num = len(run_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output: ModelOutput = self.model.forward(model_input) - b_has_out_cpu = model_input.b_prefill_has_output_cpu[0:req_num] + b_has_out_cpu = model_input.b_prefill_has_output_cpu self._capture_prompt_logprobs_if_needed(model_input, run_reqs, model_output.prompt_logics) - b_req_idx = model_input.b_req_idx[0:req_num] - b_mtp_index = model_input.b_mtp_index[0:req_num] + b_req_idx = model_input.b_req_idx + b_mtp_index = model_input.b_mtp_index if req_num > 0: ( @@ -423,7 +421,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] next_token_logprobs_cpu, next_token_ranks_cpu, ) = self._sample_and_scatter_token( - logits=model_output.logits[0:req_num, :], + logits=model_output.logits, b_req_idx=b_req_idx, b_mtp_index=b_mtp_index, run_reqs=run_reqs, @@ -660,27 +658,13 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I dtype=logits0.dtype, device=logits0.device, ) - logits[0:req_num0, :].copy_(logits0[0:req_num0, :], non_blocking=True) - logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1[0:req_num1, :], non_blocking=True) + logits[0:req_num0, :].copy_(logits0, non_blocking=True) + logits[req_num0 : (req_num0 + req_num1), :].copy_(logits1, non_blocking=True) run_reqs = run_reqs0 + run_reqs1 - b_has_out_cpu = ( - model_input0.b_prefill_has_output_cpu[0:req_num0] + model_input1.b_prefill_has_output_cpu[0:req_num1] - ) - b_mtp_index = torch.cat( - ( - model_input0.b_mtp_index[0:req_num0], - model_input1.b_mtp_index[0:req_num1], - ), - dim=0, - ) - b_req_idx = torch.cat( - ( - model_input0.b_req_idx[0:req_num0], - model_input1.b_req_idx[0:req_num1], - ), - dim=0, - ) + b_has_out_cpu = model_input0.b_prefill_has_output_cpu + model_input1.b_prefill_has_output_cpu + b_mtp_index = torch.cat((model_input0.b_mtp_index, model_input1.b_mtp_index), dim=0) + b_req_idx = torch.cat((model_input0.b_req_idx, model_input1.b_req_idx), dim=0) if req_num > 0: ( @@ -712,8 +696,7 @@ def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[I ) if req_num > 0 and g_infer_context.is_linear_att_mixed_model: - _b_req_idx = torch.cat((model_input0.b_req_idx[0:req_num0], model_input1.b_req_idx[0:req_num1]), dim=0) - g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=_b_req_idx, reqs=run_reqs) + g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=b_req_idx, reqs=run_reqs) sync_event = torch.cuda.Event() sync_event.record() From b3eb120e8cf41aa3661455f1b016b353b4b7348a Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 14:53:28 +0000 Subject: [PATCH 098/103] Require MTP row indexes during verification --- .../router/model_infer/mode_backend/dp_backend/impl.py | 2 +- lightllm/server/router/model_infer/mtp_speculative/utils.py | 5 ++--- 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 09a2d96c92..051b6c0f93 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -802,7 +802,7 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf next_token_ids=next_token_ids, b_req_idx=b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, - b_mtp_index=b_mtp_index if self.is_linear_att_mixed_model else None, + b_mtp_index=b_mtp_index, ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index d0b17274ff..935fe92d33 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections import Counter -from typing import TYPE_CHECKING, List, Optional, Tuple +from typing import TYPE_CHECKING, List, Tuple import torch @@ -39,7 +39,7 @@ def verify_mtp_tokens( next_token_ids: torch.Tensor, b_req_idx: torch.Tensor, b_req_mtp_start_loc: torch.Tensor, - b_mtp_index: Optional[torch.Tensor] = None, + b_mtp_index: torch.Tensor, ) -> Tuple[torch.Tensor, torch.Tensor]: """Verify target tokens and update recurrent MTP state when required.""" @@ -50,7 +50,6 @@ def verify_mtp_tokens( b_req_idx=b_req_idx, ) if backend.is_linear_att_mixed_model: - assert b_mtp_index is not None linear_att_mtp_state_index_update( req_to_mtp_state_index=backend.model.req_manager.req_to_mtp_state_index, b_req_mtp_start_loc=b_req_mtp_start_loc, From f5cee6b807faf264d0f93bb12fd1ec3f10c5f680 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 14:58:11 +0000 Subject: [PATCH 099/103] Return decode request splits from overlap prep --- .../mode_backend/dp_backend/impl.py | 9 ++++---- .../mode_backend/generic_pre_process.py | 8 ++++--- .../test_dp_overlap_spec_engine.py | 2 +- .../mode_backend/test_generic_pre_process.py | 22 +++++++++++++++++-- 4 files changed, 31 insertions(+), 10 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 051b6c0f93..0f68e61439 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -349,7 +349,7 @@ def prefill_overlap(self, event_pack: OverlapEventPack, prefill_reqs: List[Infer return def decode_overlap(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): - model_input0, run_reqs0, model_input1, run_reqs1 = overlap_prepare_decode_inputs(req_objs=decode_reqs) + model_input0, run_reqs0, _, model_input1, run_reqs1, _ = overlap_prepare_decode_inputs(req_objs=decode_reqs) run_reqs = run_reqs0 + run_reqs1 req_num0, req_num1 = len(run_reqs0), len(run_reqs1) @@ -728,12 +728,13 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf ( model_input0, run_reqs0, + decode_reqs0, model_input1, run_reqs1, + decode_reqs1, ) = overlap_prepare_decode_inputs(req_objs=decode_reqs) - request_split = (len(decode_reqs) + 1) // 2 - real_request_num0 = request_split - real_request_num1 = len(decode_reqs) - request_split + real_request_num0 = len(decode_reqs0) + real_request_num1 = len(decode_reqs1) req_num = real_request_num0 + real_request_num1 spec_engine = self.decode_draft_engine with torch.cuda.stream(g_infer_context.get_overlap_stream()): diff --git a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py index 5e73e11f58..22731439c4 100644 --- a/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py +++ b/lightllm/server/router/model_infer/mode_backend/generic_pre_process.py @@ -168,13 +168,15 @@ def overlap_prepare_decode_inputs(req_objs: List[InferReq]): """按请求把 decode batch 拆成两个允许为空的 microbatch。""" split_req_bound = (len(req_objs) + 1) // 2 + decode_reqs0 = req_objs[:split_req_bound] + decode_reqs1 = req_objs[split_req_bound:] model_input0, run_reqs0 = prepare_decode_inputs( - req_objs=req_objs[:split_req_bound], + req_objs=decode_reqs0, ) model_input1, run_reqs1 = prepare_decode_inputs( - req_objs=req_objs[split_req_bound:], + req_objs=decode_reqs1, ) - return model_input0, run_reqs0, model_input1, run_reqs1 + return model_input0, run_reqs0, decode_reqs0, model_input1, run_reqs1, decode_reqs1 def overlap_prepare_prefill_inputs(req_objs: List[InferReq]): diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 66543e68c9..28c53838b1 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -462,7 +462,7 @@ def propose_next_overlap(self, **kwargs): monkeypatch.setattr( dp_backend_impl, "overlap_prepare_decode_inputs", - lambda req_objs: (model_input0, [], model_input1, []), + lambda req_objs: (model_input0, [], [], model_input1, [], []), ) monkeypatch.setattr( dp_backend_impl, diff --git a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py index d22c380e19..62296634a9 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -130,8 +130,17 @@ def test_overlap_decode_builds_two_unpadded_inputs(monkeypatch): _patch_overlap_input_context(monkeypatch) reqs = [_make_decode_req(req_idx=index) for index in range(3)] - model_input0, run_reqs0, model_input1, run_reqs1 = generic_pre_process.overlap_prepare_decode_inputs(reqs) + ( + model_input0, + run_reqs0, + decode_reqs0, + model_input1, + run_reqs1, + decode_reqs1, + ) = generic_pre_process.overlap_prepare_decode_inputs(reqs) + assert decode_reqs0 == reqs[:2] + assert decode_reqs1 == reqs[2:] assert run_reqs0 == reqs[:2] assert run_reqs1 == reqs[2:] assert model_input0.batch_size == 2 @@ -144,8 +153,17 @@ def test_overlap_decode_preserves_empty_microbatch(monkeypatch): _patch_overlap_input_context(monkeypatch) req = _make_decode_req(req_idx=7) - model_input0, run_reqs0, model_input1, run_reqs1 = generic_pre_process.overlap_prepare_decode_inputs([req]) + ( + model_input0, + run_reqs0, + decode_reqs0, + model_input1, + run_reqs1, + decode_reqs1, + ) = generic_pre_process.overlap_prepare_decode_inputs([req]) + assert decode_reqs0 == [req] + assert decode_reqs1 == [] assert run_reqs0 == [req] assert model_input0.batch_size == 1 assert run_reqs1 == [] From 6514221e0247bccc4cc403d02f83e4609304a008 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 15:21:43 +0000 Subject: [PATCH 100/103] Split DP overlap MTP acceptance state --- .../mode_backend/dp_backend/impl.py | 46 +++++++++++++------ .../mtp_speculative/dp_overlap_engine.py | 13 +++--- .../dp_overlap_proposers/base.py | 5 +- .../dp_overlap_proposers/eagle_no_att.py | 13 ++---- .../dp_overlap_proposers/eagle_with_att.py | 13 ++---- .../dp_overlap_proposers/vanilla_no_att.py | 13 ++---- .../dp_overlap_proposers/vanilla_with_att.py | 17 +++---- .../test_dp_overlap_spec_engine.py | 21 ++++----- .../mtp_speculative/test_eagle_overlap.py | 25 ++++------ .../mtp_speculative/test_vanilla_overlap.py | 15 +++--- 10 files changed, 85 insertions(+), 96 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index 0f68e61439..9a81927bc1 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -805,6 +805,8 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf b_req_mtp_start_loc=b_req_mtp_start_loc, b_mtp_index=b_mtp_index, ) + mtp_accept_len0 = mtp_accept_len[:real_request_num0] + mtp_accept_len1 = mtp_accept_len[real_request_num0:] accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", gpu_tensor=accepted_index, @@ -813,9 +815,15 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf key="mtp_accept_len", gpu_tensor=mtp_accept_len, ) + accepted_index_cpu0 = accepted_index_cpu[:verify_row_num0] + accepted_index_cpu1 = accepted_index_cpu[verify_row_num0:] + mtp_accept_len_cpu0 = mtp_accept_len_cpu[:real_request_num0] + mtp_accept_len_cpu1 = mtp_accept_len_cpu[real_request_num0:] else: b_req_idx = torch.empty((0,), dtype=torch.int32, device=model_input0.b_req_idx.device) mtp_accept_len = torch.empty((0,), dtype=torch.int32, device=model_input0.b_req_idx.device) + mtp_accept_len0 = mtp_accept_len + mtp_accept_len1 = mtp_accept_len b_req_mtp_start_loc = torch.empty((0,), dtype=torch.int32, device=model_input0.b_req_idx.device) next_token_ids = torch.empty((0,), dtype=torch.int64, device=model_input0.b_req_idx.device) verify_event = torch.cuda.Event() @@ -827,12 +835,11 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf target_model_input0=model_input0, target_model_output0=model_output0, target_next_token_ids0=target_next_token_ids0, + accept_len0=mtp_accept_len0, target_model_input1=model_input1, target_model_output1=model_output1, target_next_token_ids1=target_next_token_ids1, - real_request_num0=real_request_num0, - real_request_num1=real_request_num1, - accept_len=mtp_accept_len, + accept_len1=mtp_accept_len1, draft_step=spec_plan.draft_step, ) if req_num > 0: @@ -859,11 +866,19 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf verify_event.synchronize() mtp_utils.record_request_mtp_metrics( backend=self, - decode_reqs=decode_reqs, - accept_lengths_cpu=mtp_accept_len_cpu, - verify_run_reqs=run_reqs, + decode_reqs=decode_reqs0, + accept_lengths_cpu=mtp_accept_len_cpu0, + verify_run_reqs=run_reqs0, + ) + mtp_utils.record_request_mtp_metrics( + backend=self, + decode_reqs=decode_reqs1, + accept_lengths_cpu=mtp_accept_len_cpu1, + verify_run_reqs=run_reqs1, ) - verify_ok_reqs = [req for req, accepted in zip(run_reqs, accepted_index_cpu.tolist()) if accepted] + verify_ok_reqs0 = [req for req, accepted in zip(run_reqs0, accepted_index_cpu0.tolist()) if accepted] + verify_ok_reqs1 = [req for req, accepted in zip(run_reqs1, accepted_index_cpu1.tolist()) if accepted] + verify_ok_reqs = verify_ok_reqs0 + verify_ok_reqs1 update_packs = self._pre_post_handle(verify_ok_reqs, is_chuncked_mode=False) event_pack.notify_forward_and_wait_post_handle() @@ -874,12 +889,17 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf req_num=req_num, accept_lengths_cpu=mtp_accept_len_cpu, ) - mem_indexes_cpu = torch.cat((model_input0.mem_indexes_cpu, model_input1.mem_indexes_cpu), dim=0) - proposal.extra_mem_indexes_cpu.append( - MtpMemIndexesToFree( - mem_indexes_cpu=mem_indexes_cpu, - free_mask_cpu=accepted_index_cpu == 0, - ), + proposal.extra_mem_indexes_cpu.extend( + ( + MtpMemIndexesToFree( + mem_indexes_cpu=model_input0.mem_indexes_cpu, + free_mask_cpu=accepted_index_cpu0 == 0, + ), + MtpMemIndexesToFree( + mem_indexes_cpu=model_input1.mem_indexes_cpu, + free_mask_cpu=accepted_index_cpu1 == 0, + ), + ) ) select_mask = accepted_index_cpu.to(dtype=torch.bool) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index 054e21a5ab..5200db6e16 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -144,28 +144,27 @@ def propose_next_overlap( target_model_input0: ModelInput, # batch_size = verify_batch_size0 target_model_output0: ModelOutput, # logits: [verify_batch_size0, vocab_size] target_next_token_ids0: torch.Tensor, # [verify_batch_size0] + accept_len0: torch.Tensor, # [real_req_num0] target_model_input1: ModelInput, # batch_size = verify_batch_size1 target_model_output1: ModelOutput, # logits: [verify_batch_size1, vocab_size] target_next_token_ids1: torch.Tensor, # [verify_batch_size1] - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, # [real_req_num0 + real_req_num1] + accept_len1: torch.Tensor, # [real_req_num1] draft_step: int, ) -> SpecProposal: assert target_next_token_ids0.shape == (target_model_input0.batch_size,) assert target_next_token_ids1.shape == (target_model_input1.batch_size,) - assert accept_len.shape == (real_request_num0 + real_request_num1,) + assert accept_len0.ndim == 1 + assert accept_len1.ndim == 1 return self.proposer.propose_next_overlap( target_model_input0=target_model_input0, target_model_output0=target_model_output0, target_next_token_ids0=target_next_token_ids0, + accept_len0=accept_len0, target_model_input1=target_model_input1, target_model_output1=target_model_output1, target_next_token_ids1=target_next_token_ids1, - real_request_num0=real_request_num0, - real_request_num1=real_request_num1, - accept_len=accept_len, + accept_len1=accept_len1, draft_step=draft_step, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py index bdc5c757fe..95286c7b70 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -39,12 +39,11 @@ def propose_next_overlap( target_model_input0: ModelInput, # batch_size = padded_verify_batch_size0 target_model_output0: ModelOutput, # logits: [padded_verify_batch_size0, vocab_size] target_next_token_ids0: torch.Tensor, # [padded_verify_batch_size0] + accept_len0: torch.Tensor, # [real_req_num0] target_model_input1: ModelInput, # batch_size = padded_verify_batch_size1 target_model_output1: ModelOutput, # logits: [padded_verify_batch_size1, vocab_size] target_next_token_ids1: torch.Tensor, # [padded_verify_batch_size1] - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, # [real_req_num0 + real_req_num1] + accept_len1: torch.Tensor, # [real_req_num1] draft_step: int, ) -> SpecProposal: """Generate one proposal from two DP-overlapped decode microbatches.""" diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py index 2843ba496b..077d478fdc 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -34,17 +34,16 @@ def propose_next_overlap( target_model_input0: ModelInput, target_model_output0: ModelOutput, target_next_token_ids0: torch.Tensor, + accept_len0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, target_next_token_ids1: torch.Tensor, - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, + accept_len1: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: - req_num_by_batch = (int(real_request_num0), int(real_request_num1)) + accept_len_by_batch = (accept_len0, accept_len1) + req_num_by_batch = (accept_len0.shape[0], accept_len1.shape[0]) req_num = sum(req_num_by_batch) - assert accept_len.shape == (req_num,) proposal_token_ids = target_next_token_ids0.new_empty((req_num, draft_step)) schedule_scores = ( @@ -66,10 +65,6 @@ def propose_next_overlap( assert target_next_token_ids0.shape == (target_model_input0.batch_size,) assert target_next_token_ids1.shape == (target_model_input1.batch_size,) - accept_len_by_batch = ( - accept_len[: req_num_by_batch[0]], - accept_len[req_num_by_batch[0] :], - ) b_req_mtp_start_loc = ( get_dp_overlap_req_start_rows( b_mtp_index=target_model_input0.b_mtp_index, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index d582fe7143..6b8c23e8fd 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -57,12 +57,11 @@ def propose_next_overlap( target_model_input0: ModelInput, target_model_output0: ModelOutput, target_next_token_ids0: torch.Tensor, + accept_len0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, target_next_token_ids1: torch.Tensor, - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, + accept_len1: torch.Tensor, draft_step: int, ) -> EagleSpecProposal: """提交两个 target verify microbatch 的 draft KV,并生成下一轮 proposal。""" @@ -72,9 +71,9 @@ def propose_next_overlap( assert not target_model_input1.is_prefill assert len(self.backend.draft_models) == 1 - req_num_by_batch = (int(real_request_num0), int(real_request_num1)) + accept_len_by_batch = (accept_len0, accept_len1) + req_num_by_batch = (accept_len0.shape[0], accept_len1.shape[0]) req_num = sum(req_num_by_batch) - assert accept_len.shape == (req_num,) assert target_next_token_ids0.shape == (target_model_input0.batch_size,) assert target_next_token_ids1.shape == (target_model_input1.batch_size,) assert target_model_output0.mtp_collector.spec_hidden.shape[0] == target_model_input0.batch_size @@ -82,10 +81,6 @@ def propose_next_overlap( assert target_model_input0.b_position_delta is not None assert target_model_input1.b_position_delta is not None - accept_len_by_batch = ( - accept_len[: req_num_by_batch[0]], - accept_len[req_num_by_batch[0] :], - ) b_req_mtp_start_loc = ( get_dp_overlap_req_start_rows( b_mtp_index=target_model_input0.b_mtp_index, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py index ca003c7ab4..4c3261aa26 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -34,17 +34,16 @@ def propose_next_overlap( target_model_input0: ModelInput, target_model_output0: ModelOutput, target_next_token_ids0: torch.Tensor, + accept_len0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, target_next_token_ids1: torch.Tensor, - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, + accept_len1: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: - req_num_by_batch = (int(real_request_num0), int(real_request_num1)) + accept_len_by_batch = (accept_len0, accept_len1) + req_num_by_batch = (accept_len0.shape[0], accept_len1.shape[0]) req_num = sum(req_num_by_batch) - assert accept_len.shape == (req_num,) proposal_token_ids = target_next_token_ids0.new_empty((req_num, draft_step)) schedule_scores = ( @@ -66,10 +65,6 @@ def propose_next_overlap( assert target_next_token_ids0.shape == (target_model_input0.batch_size,) assert target_next_token_ids1.shape == (target_model_input1.batch_size,) - accept_len_by_batch = ( - accept_len[: req_num_by_batch[0]], - accept_len[req_num_by_batch[0] :], - ) b_req_mtp_start_loc = ( get_dp_overlap_req_start_rows( b_mtp_index=target_model_input0.b_mtp_index, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py index ecda2dfb21..4f3c117b34 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -64,36 +64,31 @@ def propose_next_overlap( target_model_input0: ModelInput, target_model_output0: ModelOutput, target_next_token_ids0: torch.Tensor, + accept_len0: torch.Tensor, target_model_input1: ModelInput, target_model_output1: ModelOutput, target_next_token_ids1: torch.Tensor, - real_request_num0: int, - real_request_num1: int, - accept_len: torch.Tensor, + accept_len1: torch.Tensor, draft_step: int, ) -> VanillaSpecProposal: assert draft_step == self.backend.max_draft_step assert len(self.backend.draft_models) == draft_step - req_num_by_batch = (int(real_request_num0), int(real_request_num1)) - assert accept_len.shape == (sum(req_num_by_batch),) + accept_len_by_batch = (accept_len0, accept_len1) + req_num_by_batch = (accept_len0.shape[0], accept_len1.shape[0]) assert target_next_token_ids0.shape == (target_model_input0.batch_size,) assert target_next_token_ids1.shape == (target_model_input1.batch_size,) b_req_mtp_start_loc = ( get_dp_overlap_req_start_rows( b_mtp_index=target_model_input0.b_mtp_index, - req_num=real_request_num0, + req_num=req_num_by_batch[0], ), get_dp_overlap_req_start_rows( b_mtp_index=target_model_input1.b_mtp_index, - req_num=real_request_num1, + req_num=req_num_by_batch[1], ), ) - accept_len_by_batch = ( - accept_len[: req_num_by_batch[0]], - accept_len[req_num_by_batch[0] : sum(req_num_by_batch)], - ) accepted_tail_rows = tuple( req_mtp_start_loc + batch_accept_len - 1 for req_mtp_start_loc, batch_accept_len in zip( diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 28c53838b1..fcdf4e3295 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -312,18 +312,18 @@ def propose_next_overlap(self, **kwargs): model_output1 = SimpleNamespace() target_next_token_ids0 = torch.arange(8, dtype=torch.int64) target_next_token_ids1 = torch.arange(8, 24, dtype=torch.int64) - accept_len = torch.tensor([2, 3, 4], dtype=torch.int32) + accept_len0 = torch.tensor([2], dtype=torch.int32) + accept_len1 = torch.tensor([3, 4], dtype=torch.int32) proposal = engine.propose_next_overlap( target_model_input0=model_input0, target_model_output0=model_output0, target_next_token_ids0=target_next_token_ids0, + accept_len0=accept_len0, target_model_input1=model_input1, target_model_output1=model_output1, target_next_token_ids1=target_next_token_ids1, - real_request_num0=1, - real_request_num1=2, - accept_len=accept_len, + accept_len1=accept_len1, draft_step=7, ) @@ -333,9 +333,8 @@ def propose_next_overlap(self, **kwargs): assert calls["target_model_output1"] is model_output1 assert calls["target_next_token_ids0"] is target_next_token_ids0 assert calls["target_next_token_ids1"] is target_next_token_ids1 - assert calls["real_request_num0"] == 1 - assert calls["real_request_num1"] == 2 - assert calls["accept_len"] is accept_len + assert calls["accept_len0"] is accept_len0 + assert calls["accept_len1"] is accept_len1 assert calls["draft_step"] == 7 assert proposal.token_ids.shape == (3, 7) @@ -442,10 +441,10 @@ def propose_next_overlap(self, **kwargs): assert kwargs["target_next_token_ids0"].dtype == torch.int64 assert kwargs["target_next_token_ids1"].shape == (0,) assert kwargs["target_next_token_ids1"].dtype == torch.int64 - assert kwargs["real_request_num0"] == 0 - assert kwargs["real_request_num1"] == 0 - assert kwargs["accept_len"].shape == (0,) - assert kwargs["accept_len"].dtype == torch.int32 + assert kwargs["accept_len0"].shape == (0,) + assert kwargs["accept_len0"].dtype == torch.int32 + assert kwargs["accept_len1"].shape == (0,) + assert kwargs["accept_len1"].dtype == torch.int32 assert kwargs["draft_step"] == 2 return SpecProposal(token_ids=torch.empty((0, 2), dtype=torch.int64, device=device)) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 6e35a37185..2e24262e7e 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -133,15 +133,14 @@ def test_overlap_eagle_supports_variable_verify_layout(monkeypatch): mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((3, 2))), ), target_next_token_ids0=torch.arange(3, dtype=torch.int64), + accept_len0=torch.tensor([2], dtype=torch.int32), target_model_input1=model_input1, target_model_output1=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), ), target_next_token_ids1=torch.arange(10, 16, dtype=torch.int64), - real_request_num0=1, - real_request_num1=2, - accept_len=torch.tensor([2, 1, 3], dtype=torch.int32), + accept_len1=torch.tensor([1, 3], dtype=torch.int32), draft_step=2, ) @@ -185,15 +184,14 @@ def test_overlap_eagle_supports_empty_verify_rows(monkeypatch): mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((0, 2))), ), target_next_token_ids0=torch.empty((0,), dtype=torch.int64), + accept_len0=torch.empty((0,), dtype=torch.int32), target_model_input1=_target_input(batch_size=0), target_model_output1=ModelOutput( logits=torch.empty((0, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((0, 2))), ), target_next_token_ids1=torch.empty((0,), dtype=torch.int64), - real_request_num0=0, - real_request_num1=0, - accept_len=torch.empty((0,), dtype=torch.int32), + accept_len1=torch.empty((0,), dtype=torch.int32), draft_step=2, ) @@ -234,6 +232,7 @@ def test_overlap_eagle_returns_dynamic_schedule_scores(monkeypatch): mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((2, 2))), ), target_next_token_ids0=torch.arange(2, dtype=torch.int64), + accept_len0=torch.tensor([2], dtype=torch.int32), target_model_input1=_target_input( batch_size=3, b_mtp_index=torch.tensor([0, 1, 0], dtype=torch.int32), @@ -243,9 +242,7 @@ def test_overlap_eagle_returns_dynamic_schedule_scores(monkeypatch): mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((3, 2))), ), target_next_token_ids1=torch.arange(2, 5, dtype=torch.int64), - real_request_num0=1, - real_request_num1=2, - accept_len=torch.tensor([2, 1, 1], dtype=torch.int32), + accept_len1=torch.tensor([1, 1], dtype=torch.int32), draft_step=2, ) @@ -287,15 +284,14 @@ def test_overlap_eagle_no_att_supports_dynamic_draft_step(): mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((2, 2), device=device)), ), target_next_token_ids0=torch.arange(2, dtype=torch.int64, device=device), + accept_len0=torch.tensor([2], dtype=torch.int32, device=device), target_model_input1=model_input1, target_model_output1=ModelOutput( logits=torch.empty((3, 1), device=device), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((3, 2), device=device)), ), target_next_token_ids1=torch.arange(2, 5, dtype=torch.int64, device=device), - real_request_num0=1, - real_request_num1=2, - accept_len=torch.tensor([2, 1, 1], dtype=torch.int32, device=device), + accept_len1=torch.tensor([1, 1], dtype=torch.int32, device=device), draft_step=2, ) @@ -339,15 +335,14 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch): mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((3, 2))), ), target_next_token_ids0=torch.arange(3, dtype=torch.int64), + accept_len0=torch.tensor([2], dtype=torch.int32), target_model_input1=model_input1, target_model_output1=ModelOutput( logits=torch.empty((6, 1)), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2))), ), target_next_token_ids1=torch.arange(10, 16, dtype=torch.int64), - real_request_num0=1, - real_request_num1=2, - accept_len=torch.tensor([2, 1, 3], dtype=torch.int32), + accept_len1=torch.tensor([1, 3], dtype=torch.int32), draft_step=2, ) diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py index 1a40d33448..b98152c855 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -48,12 +48,11 @@ def test_dp_vanilla_no_att_supports_zero_dynamic_draft_step(): target_model_input0=None, target_model_output0=None, target_next_token_ids0=target_next_token_ids0, + accept_len0=torch.ones((1,), dtype=torch.int32), target_model_input1=None, target_model_output1=None, target_next_token_ids1=target_next_token_ids1, - real_request_num0=1, - real_request_num1=1, - accept_len=torch.ones((2,), dtype=torch.int32), + accept_len1=torch.ones((1,), dtype=torch.int32), draft_step=0, ) @@ -90,15 +89,14 @@ def test_dp_vanilla_proposer_owns_overlap_decode(): mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), ), target_next_token_ids0=torch.tensor([10, 11, 0], dtype=torch.int64, device=device), + accept_len0=torch.tensor([2], dtype=torch.int32, device=device), target_model_input1=model_input1, target_model_output1=ModelOutput( logits=torch.empty((6, 1), device=device), mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((6, 2), device=device)), ), target_next_token_ids1=torch.tensor([20, 21, 22], dtype=torch.int64, device=device), - real_request_num0=1, - real_request_num1=1, - accept_len=torch.tensor([2, 1], dtype=torch.int32, device=device), + accept_len1=torch.tensor([1], dtype=torch.int32, device=device), draft_step=2, ) @@ -137,12 +135,11 @@ def test_dp_vanilla_proposer_builds_padded_inputs_for_empty_verify_rows(): target_model_input0=model_input0, target_model_output0=model_output0, target_next_token_ids0=torch.empty((0,), dtype=torch.int64, device=device), + accept_len0=torch.empty((0,), dtype=torch.int32, device=device), target_model_input1=model_input1, target_model_output1=model_output1, target_next_token_ids1=torch.empty((0,), dtype=torch.int64, device=device), - real_request_num0=0, - real_request_num1=0, - accept_len=torch.empty((0,), dtype=torch.int32, device=device), + accept_len1=torch.empty((0,), dtype=torch.int32, device=device), draft_step=2, ) From 327ed2b2f461d1299a7c71c5276fde09f3eb6b69 Mon Sep 17 00:00:00 2001 From: wangzaijun Date: Sun, 23 Aug 2026 15:41:04 +0000 Subject: [PATCH 101/103] Align MTP memory reservation across ranks --- lightllm/utils/profile_max_tokens.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/lightllm/utils/profile_max_tokens.py b/lightllm/utils/profile_max_tokens.py index eb576a966f..5ec2a9145e 100644 --- a/lightllm/utils/profile_max_tokens.py +++ b/lightllm/utils/profile_max_tokens.py @@ -42,6 +42,21 @@ def get_mtp_adjusted_mem_fraction( mtp_weight_bytes = target_weight_bytes * mtp_layer_num / target_layer_num total_gpu_bytes = torch.cuda.get_device_properties(get_current_device_id()).total_memory adjusted_mem_fraction = mem_fraction - mtp_weight_bytes / total_gpu_bytes + + # 不同 rank 的权重分片大小和 GPU 总显存可能不同。取全局最小值,保证所有 + # rank 使用相同且能够安全预留 MTP 权重显存的 KV cache 比例。 + if torch.distributed.is_initialized() and torch.distributed.get_world_size() > 1: + adjusted_mem_fraction_tensor = torch.tensor( + adjusted_mem_fraction, + dtype=torch.float32, + device=f"cuda:{get_current_device_id()}", + ) + torch.distributed.all_reduce( + adjusted_mem_fraction_tensor, + op=torch.distributed.ReduceOp.MIN, + ) + adjusted_mem_fraction = adjusted_mem_fraction_tensor.item() + assert adjusted_mem_fraction > 0, "MTP weight reservation leaves no memory for the target model KV cache" return adjusted_mem_fraction From 86f72393bbf50467ef8393701cb7c6e098bbd446 Mon Sep 17 00:00:00 2001 From: baishhao Date: Mon, 24 Aug 2026 14:08:58 +0800 Subject: [PATCH 102/103] fix: share DFlash hidden KV commit path --- lightllm/models/qwen3_dflash/model.py | 27 ++++++++++ lightllm/models/qwen3_dspark/model.py | 28 ---------- .../models/test_qwen3_dspark_model_output.py | 54 +++++++++++++++++++ 3 files changed, 81 insertions(+), 28 deletions(-) diff --git a/lightllm/models/qwen3_dflash/model.py b/lightllm/models/qwen3_dflash/model.py index 48106ac730..122f4044c3 100644 --- a/lightllm/models/qwen3_dflash/model.py +++ b/lightllm/models/qwen3_dflash/model.py @@ -1,8 +1,11 @@ +import torch + from lightllm.common.basemodel.attention import ( Fa3AttBackend, Fp8Fa3AttBackend, ) from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput from lightllm.models.llama.model import LlamaTpPartModel from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_dflash.infer_struct import Qwen3DFlashInferStateInfo @@ -81,6 +84,30 @@ def _init_weights(self, start_layer_index=None): for i in range(self.config["n_layer"]) ] + def _decode(self, model_input: ModelInput) -> ModelOutput: + if model_input.mtp_draft_input_hiddens is None: + return super()._decode(model_input) + + assert model_input.mtp_draft_input_hiddens.shape[0] == model_input.batch_size + + # Target verification already computed the hidden rows that need to be + # committed to the draft cache. Project those rows and write their KV + # directly: no token embedding, attention state, or draft logits are + # needed for this half of the parallel-block proposal. + position_ids = model_input.b_seq_len - 1 + infer_state = self.infer_state_class() + infer_state.mtp_draft_input_hiddens = model_input.mtp_draft_input_hiddens + infer_state.position_cos = torch.index_select(self._cos_cached, 0, position_ids) + infer_state.position_sin = torch.index_select(self._sin_cached, 0, position_ids) + infer_state.mem_manager = self.mem_manager + infer_state.mem_index = model_input.mem_indexes + + hidden = self.pre_infer.context_forward(None, infer_state, self.pre_post_weight) + for layer, layer_weight in zip(self.layers_infer, self.trans_layers_weight): + hidden = layer.context_forward(hidden, infer_state, layer_weight) + + return ModelOutput(logits=hidden.new_empty((model_input.batch_size, 1))) + def _gen_special_model_input(self, token_num: int): return {"mtp_draft_input_hiddens": None} diff --git a/lightllm/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py index aba99c9851..a9e49ee91f 100644 --- a/lightllm/models/qwen3_dspark/model.py +++ b/lightllm/models/qwen3_dspark/model.py @@ -1,8 +1,3 @@ -from types import SimpleNamespace - -import torch - -from lightllm.common.basemodel.batch_objs import ModelOutput from lightllm.models.draft_registry import DraftModelRegistry from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import Qwen3DSparkPostLayerInfer @@ -20,26 +15,3 @@ class Qwen3DSparkModel(Qwen3DFlashModel): pre_and_post_weight_class = Qwen3DSparkPreAndPostLayerWeight post_layer_infer_class = Qwen3DSparkPostLayerInfer - - def _decode(self, model_input): - if model_input.mtp_draft_input_hiddens is None: - return super()._decode(model_input) - - assert model_input.mtp_draft_input_hiddens.shape[0] == model_input.batch_size - - # This path only projects target hidden and writes its RoPE'd KV. It - # deliberately avoids the full decode infer state and attention setup. - position_ids = model_input.b_seq_len - 1 - infer_state = SimpleNamespace( - mtp_draft_input_hiddens=model_input.mtp_draft_input_hiddens, - position_cos=torch.index_select(self._cos_cached, 0, position_ids), - position_sin=torch.index_select(self._sin_cached, 0, position_ids), - mem_manager=self.mem_manager, - mem_index=model_input.mem_indexes, - ) - - hidden = self.pre_infer.context_forward(None, infer_state, self.pre_post_weight) - for layer, layer_weight in zip(self.layers_infer, self.trans_layers_weight): - hidden = layer.context_forward(hidden, infer_state, layer_weight) - - return ModelOutput(logits=hidden.new_empty((model_input.batch_size, 1))) diff --git a/unit_tests/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py index 6be526d8bb..1dfdfeae23 100644 --- a/unit_tests/models/test_qwen3_dspark_model_output.py +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -5,6 +5,7 @@ from lightllm.common.basemodel import batch_objs from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput +from lightllm.models.qwen3_5_dflash.model import Qwen3_5DFlashModel from lightllm.models.qwen3_5_dspark.model import Qwen3_5DSparkModel from lightllm.models.qwen3_dflash.layer_infer import transformer_layer_infer as qwen3_dflash_layer_infer from lightllm.models.qwen3_dflash.layer_infer.transformer_layer_infer import Qwen3DFlashTransformerLayerInfer @@ -14,6 +15,59 @@ from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel +@pytest.mark.parametrize( + "model_class", + [Qwen3DFlashModel, Qwen3_5DFlashModel, Qwen3DSparkModel, Qwen3_5DSparkModel], +) +def test_parallel_block_decode_commits_target_hiddens_directly(model_class): + model = model_class.__new__(model_class) + model._cos_cached = torch.arange(24).view(6, 4) + model._sin_cached = model._cos_cached + 100 + model.mem_manager = object() + model.pre_post_weight = object() + + target_hiddens = torch.arange(6, dtype=torch.float32).view(2, 3) + mem_indexes = torch.tensor([7, 11]) + observed_states = [] + + class PreInfer: + def context_forward(self, input_ids, infer_state, layer_weight): + assert input_ids is None + assert layer_weight is model.pre_post_weight + assert infer_state.mtp_draft_input_hiddens is target_hiddens + observed_states.append(infer_state) + return target_hiddens + 1 + + class TransformerLayer: + def __init__(self, increment): + self.increment = increment + + def context_forward(self, hidden, infer_state, layer_weight): + assert infer_state is observed_states[0] + assert layer_weight == self.increment + return hidden + self.increment + + model.pre_infer = PreInfer() + model.layers_infer = [TransformerLayer(2), TransformerLayer(3)] + model.trans_layers_weight = [2, 3] + model_input = SimpleNamespace( + batch_size=2, + b_seq_len=torch.tensor([3, 5]), + mem_indexes=mem_indexes, + mtp_draft_input_hiddens=target_hiddens, + ) + + output = model._decode(model_input) + + assert isinstance(output, ModelOutput) + assert output.logits.shape == (2, 1) + infer_state = observed_states[0] + assert infer_state.mem_manager is model.mem_manager + assert infer_state.mem_index is mem_indexes + assert torch.equal(infer_state.position_cos, model._cos_cached[[2, 4]]) + assert torch.equal(infer_state.position_sin, model._sin_cached[[2, 4]]) + + @pytest.mark.parametrize("model_class", [Qwen3DSparkModel, Qwen3_5DSparkModel]) def test_dspark_decode_unpad_uses_common_output_and_slices_mtp_fields(model_class): model = model_class.__new__(model_class) From 9cbc9c04d32d9649ce89a8a8b325666a6caacf89 Mon Sep 17 00:00:00 2001 From: baishhao Date: Mon, 24 Aug 2026 14:09:35 +0800 Subject: [PATCH 103/103] chore: update Qwen3.5 DSpark 1P1D launch --- test/start_scripts/qwen35/qwen35_pd_1p1d.sh | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/test/start_scripts/qwen35/qwen35_pd_1p1d.sh b/test/start_scripts/qwen35/qwen35_pd_1p1d.sh index 35e76780bf..2a5fc17298 100755 --- a/test/start_scripts/qwen35/qwen35_pd_1p1d.sh +++ b/test/start_scripts/qwen35/qwen35_pd_1p1d.sh @@ -30,7 +30,7 @@ P_COMMON_ARGS=( --model_name qwen35_27b --graph_max_batch_size 8 --running_max_req_size 8 - --mem_fraction 0.80 + --mem_fraction 0.75 --max_image_token_count 4096 --max_image_pixels 3686400 --batch_max_tokens 8192 @@ -40,13 +40,12 @@ P_COMMON_ARGS=( --quant_type fp8w8a8-pt-sgl --mtp_mode dspark --mtp_draft_model_dir "${DSPARK_MODEL_DIR}" - --mtp_step 1 + --mtp_step 5 "${CHAT_TEMPLATE_ARGS[@]}" --pd_trans_mode nccl --pd_kv_page_size 4096 --pd_master_ip 127.0.0.1 --pd_master_port "${PORT}" - --enable_prefill_cudagraph ) D_COMMON_ARGS=(