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 diff --git a/lightllm/common/basemodel/attention/base_att.py b/lightllm/common/basemodel/attention/base_att.py index 55d97d2aa8..a7e2d8122a 100644 --- a/lightllm/common/basemodel/attention/base_att.py +++ b/lightllm/common/basemodel/attention/base_att.py @@ -1,8 +1,13 @@ +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: from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.common.basemodel.infer_struct import InferStateInfo @@ -10,33 +15,65 @@ class BaseAttBackend: """ - 用于创建支持各种不同的AttBackend, 如 fa3, flashinfer, triton 实现等, - 这个是单列模式, 每种backend只有一个实例 + 用于创建支持各种不同的AttBackend, 如 fa3, flashinfer, triton 实现等。 + 每个 model 复用一个 backend 实例。 """ _instances = {} + _workspace_buffers = {} + _workspace_buffer_lock = threading.Lock() 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, 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 + @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") def create_att_decode_state(self) -> "BaseDecodeAttState": raise NotImplementedError("not impl") + 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) + 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 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 57f3ab6fe3..0f0f62ec63 100644 --- a/lightllm/common/basemodel/attention/fa3/fp.py +++ b/lightllm/common/basemodel/attention/fa3/fp.py @@ -2,11 +2,14 @@ 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.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.envs_utils import get_env_start_args class Fa3AttBackend(BaseAttBackend): @@ -20,13 +23,20 @@ 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.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(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() - ), + self.get_gpu_workspace_buffer( + key_name=f"fa3_fp_page_table_{buffer_index}", + workspace_size=workspace_size, + dtype=torch.int32, + ) + for buffer_index in range(buffer_count) ] return self._shared_page_table_buffer @@ -42,8 +52,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( @@ -102,7 +114,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=self.causal, window_size=window_size, softcap=0.0, k_descale=k_descale, @@ -121,31 +133,71 @@ 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 - args_mtp_step = get_env_start_args().mtp_step - - if args_mtp_step > 0: - # 修正 mtp 在 fa3 下的输入。 - mtp_size = args_mtp_step + 1 - 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[mtp_size - 1 :: mtp_size] - 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() + 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) + elif draft_step > 0: + b_att_req_idx = self._init_fixed_spec_decode_state(draft_step) 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() + b_att_req_idx = self._init_normal_decode_state() - att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) - assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 + 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_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, + ) + 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中申请 if ( @@ -163,23 +215,11 @@ def init_state(self): device=self.infer_state.input_ids.device, ) - if 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 - return + 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, + ) def copy_for_decode_cuda_graph(self, new_state: "Fa3DecodeAttState"): super().copy_for_decode_cuda_graph(new_state) @@ -232,7 +272,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=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 adc8b5c01e..6e6d8a8366 100644 --- a/lightllm/common/basemodel/attention/fa3/fp8.py +++ b/lightllm/common/basemodel/attention/fa3/fp8.py @@ -3,7 +3,6 @@ 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 @@ -99,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=True, + causal=self.causal, window_size=(-1, -1), softcap=0.0, q_descale=q_scale, @@ -119,11 +118,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 + att_batch_size = self.b_att_seq_len.shape[0] mem_manager = self.backend.model.mem_manager offline_scales: torch.Tensor = mem_manager.scales @@ -131,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 @@ -190,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=True, + 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 9a10457b12..5b757a0d67 100644 --- a/lightllm/common/basemodel/attention/fa3/mla.py +++ b/lightllm/common/basemodel/attention/fa3/mla.py @@ -2,12 +2,12 @@ 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.utils.envs_utils import get_env_start_args -from lightllm.common.basemodel.triton_kernel.fa3_utils import 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.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 class MlaFa3AttBackend(BaseAttBackend): @@ -21,13 +21,20 @@ 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.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(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() - ), + self.get_gpu_workspace_buffer( + key_name=f"fa3_mla_page_table_{buffer_index}", + workspace_size=workspace_size, + dtype=torch.int32, + ) + for buffer_index in range(buffer_count) ] return self._shared_page_table_buffer @@ -42,8 +49,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() @@ -90,7 +99,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=self.causal, return_softmax_lse=False, ) return o_tensor @@ -104,31 +113,71 @@ 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 - args_mtp_step = get_env_start_args().mtp_step - - if args_mtp_step > 0: - # 修正 mtp 在 fa3 下的输入。 - mtp_size = args_mtp_step + 1 - 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[mtp_size - 1 :: mtp_size] - 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() + 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) + elif draft_step > 0: + b_att_req_idx = self._init_fixed_spec_decode_state(draft_step) 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() + b_att_req_idx = self._init_normal_decode_state() + + 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_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, + ) + 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 - att_batch_size = self.infer_state.batch_size // (args_mtp_step + 1) - assert self.infer_state.batch_size % (args_mtp_step + 1) == 0 + 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中申请 if ( @@ -146,23 +195,11 @@ def init_state(self): device=self.infer_state.input_ids.device, ) - if 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 - return + 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, + ) def copy_for_decode_cuda_graph(self, new_state: "MlaFa3DecodeAttState"): super().copy_for_decode_cuda_graph(new_state) @@ -219,7 +256,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=self.causal, window_size=(-1, -1), softcap=0.0, k_descale=k_descale, 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/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index 404ebf17cc..ca6ceaec43 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -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.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, ) @@ -203,35 +206,61 @@ class LinearAttDecodeAttState(BaseDecodeAttState): b_num_accepted_tokens: torch.Tensor = None def init_state(self): - backend: LinearAttBackend = self.backend - mtp_step = backend.mtp_step + draft_step = self.backend.model.mtp_manager.get_decode_draft_step(self.backend.model.is_mtp_draft_model) + if draft_step == 0: + 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 - # decode 模式下 - if mtp_step == 0: - # 非mtp模式下,不需要额外状态 - 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 - 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信息的起点。 - self.b_num_accepted_tokens = self.infer_state.req_manager.req_to_mtp_state_index[self.b_conv_buffer_idx] + 1 - return + 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 + 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): + 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, + ).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, @@ -253,8 +282,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 状态。 + 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_mtp_kernel( mixed_qkv, conv_states, @@ -335,18 +364,18 @@ def _gdn_mtp_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_mtp_cu_q_seq_len - mixed_qkv = causal_conv1d_update_spec( + mixed_qkv = causal_conv1d_update_mtp( mixed_qkv, conv_states, layer_weight.linear_conv1d.mm_param.weight, - mtp_step=backend.mtp_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/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..c58f88244e 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, ) @@ -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/attention/triton/fp.py b/lightllm/common/basemodel/attention/triton/fp.py index a1370a7045..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) @@ -112,9 +120,15 @@ 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.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] - if q_head_num == k_head_num: + + 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: @@ -205,6 +219,30 @@ def _normal_decode_gqa_flash_decoding_att( return out + def _spec_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, + ) + + 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=self.infer_state.b_seq_len, + b_mark_shared_group=self.b_mark_mtp_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/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 0f1bfa9cc6..f2b6bae085 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -25,13 +25,21 @@ 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 ( + NoopHiddenCollector, +) +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, - enable_diverse_mode_gqa_decode_fast_kernel, enable_full_att_decode_tune, ) from lightllm.common.triton_utils.autotuner import Autotuner @@ -49,6 +57,8 @@ class TpPartBaseModel: + is_mtp_draft_model = False + # weight class pre_and_post_weight_class = None transformer_weight_class = None @@ -78,15 +88,16 @@ 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 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.mtp_manager = MtpManager.get_instance() + 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) @@ -98,12 +109,6 @@ 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._init_config() @@ -112,8 +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): - 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() @@ -136,6 +142,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() self._autotune_warmup() self._full_att_decode_autotune() self._init_padded_req() @@ -272,13 +279,19 @@ def _init_att_backend1(self): return def _init_cudagraph(self): + 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_, + capture_infer_cost=self.args.mtp_dynamic_verify, ) ) if self.graph is not None: @@ -318,7 +331,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. @@ -329,7 +342,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 @@ -338,28 +351,40 @@ def _full_att_decode_autotune(self): from lightllm.utils.sgl_utils import fa3_decode_autotune + 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_, ) - fa3_decode_autotune(self, cuda_graph_batch_sizes) + 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): pass + def _init_hidden_collector(self): + self.hidden_collector_prototype = self.mtp_manager.create_hidden_collector(model=self) + @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) 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 @@ -369,7 +394,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 @@ -381,9 +405,8 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) 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 + infer_state.b_shared_seq_len = model_input.b_shared_seq_len + infer_state.b_shared_radix_node_id = model_input.b_shared_radix_node_id infer_state.multimodal_params = model_input.multimodal_params @@ -445,15 +468,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: @@ -466,12 +486,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 @@ -492,11 +513,11 @@ 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 可能会让外面使用的数组引用发生变化,导致错误。 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": []} @@ -518,12 +539,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] - - # 特殊模型,特殊模式的特殊变量的特殊 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] - + new_model_output.mtp_collector = model_output.mtp_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( @@ -532,23 +551,21 @@ 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.mtp_collector = padded_model_output.mtp_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: 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] - return new_model_output 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 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, @@ -558,13 +575,14 @@ 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 + # 即使当前 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( @@ -590,7 +608,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, @@ -604,62 +622,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( + 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 规范化为 + # 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): @@ -672,12 +692,20 @@ 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] + 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 [_input_embs] handle_token_num = infer_state.input_ids.shape[0] @@ -710,22 +738,21 @@ 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.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() + hidden_collector = infer_state.hidden_collector + hidden_collector.add_final_hidden(last_input_embs) + model_output = ModelOutput( + logits=predict_logits.contiguous(), + mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + 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): + 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) @@ -734,18 +761,18 @@ 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]) + 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 ) - model_output = ModelOutput(logits=predict_logits.contiguous()) - - # 特殊模型特殊模式的额外输出 - 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() + hidden_collector.add_final_hidden(last_input_embs) + model_output = ModelOutput( + logits=predict_logits.contiguous(), + mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + ) # 在 cuda graph 模式下,输出需要转为 no ref tensor, 加强mem pool 的复用,降低显存的使用。 if infer_state.is_cuda_graph: @@ -755,37 +782,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 and model_input0.b_is_decode_req is not None: - 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 and model_input1.b_is_decode_req is not None: - 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.total_token_num - model_input0.prefix_total_token_num - origin_handle_token_num1 = model_input1.total_token_num - model_input1.prefix_total_token_num - 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_handle_token_num0 = model_input0.input_ids.shape[0] + origin_handle_token_num1 = model_input1.input_ids.shape[0] + 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 @@ -846,36 +872,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) @@ -912,9 +941,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) @@ -939,14 +967,16 @@ 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 @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 @@ -967,6 +997,14 @@ 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] ) + hidden_collector0.add( + layer_index=i, + hidden=input_embs, + ) + hidden_collector1.add( + layer_index=i, + hidden=input_embs1, + ) # 折叠模式调用完infer_state 和 infer_state1 上的hook函数后,input_embs 和 input_embs1 才具备正确的运算数据。 infer_state.call_overlap_hook() @@ -983,22 +1021,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.is_mtp_mode: - 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) - 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() + hidden_collector0.add_final_hidden(last_input_embs) + hidden_collector1.add_final_hidden(last_input_embs1) + model_output = ModelOutput( + logits=predict_logits.contiguous(), + 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(), + mtp_collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), + prompt_logics=infer_state1.prompt_logics, + ) return model_output, model_output1 @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 ) @@ -1009,6 +1050,14 @@ 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] ) + hidden_collector0.add( + layer_index=i, + hidden=input_embs, + ) + hidden_collector1.add( + layer_index=i, + hidden=input_embs1, + ) # 折叠模式调用完infer_state 上的hook函数后,input_embs 和 input_embs 才具备正确的运算数据。 infer_state.call_overlap_hook() @@ -1021,14 +1070,16 @@ 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.is_mtp_mode: - 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) - model_output.mtp_main_output_hiddens = input_embs.contiguous() - model_output1.mtp_main_output_hiddens = input_embs1.contiguous() + hidden_collector0.add_final_hidden(last_input_embs) + hidden_collector1.add_final_hidden(last_input_embs1) + model_output = ModelOutput( + logits=predict_logits.contiguous(), + mtp_collector=infer_state.hidden_collector.finish_output(infer_state=infer_state), + ) + model_output1 = ModelOutput( + logits=predict_logits1.contiguous(), + mtp_collector=infer_state1.hidden_collector.finish_output(infer_state=infer_state1), + ) if infer_state.is_cuda_graph: model_output.to_no_ref_tensor() @@ -1059,21 +1110,23 @@ 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, 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, 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, + b_prefill_has_output_cpu=[False], multimodal_params=[{"images": [], "audios": []}], ) model_output = self.forward( @@ -1136,18 +1189,19 @@ 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, 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, 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, @@ -1201,18 +1255,19 @@ 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, 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, 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=[ @@ -1240,8 +1295,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" ) diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 58922b8055..ae645d4b7b 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -1,8 +1,9 @@ +import copy +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 + from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor @@ -15,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 @@ -23,23 +23,21 @@ 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 - # 只会在 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 - # 只会在继承 Qwen2VLInferStateInfo 的 MRoPE 模型 decode 阶段使用,如 - # Qwen2/2.5-VL、Qwen3-VL/MOE/Omni、Qwen3.5;普通模型不会使用。 + # 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 @@ -47,7 +45,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 # 专有变量,用于一些特殊的模型,特殊的模式下, 传递一些特殊 # 的输入变量。只在特殊的模型模式下才会具体使用和生效。 @@ -58,72 +57,153 @@ 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: + + 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) - 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." + self.b_prefill_start_loc = self.b_prefill_start_loc.cuda(non_blocking=True) + self.b_is_decode_req = self.b_is_decode_req.cuda(non_blocking=True) else: - assert self.is_prefill is True, "decode ModelInput should provide b_position_delta." + # 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) - 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) + # 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}" + 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.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" + 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 + 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: + """保存一次模型 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 + # 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. + mtp_collector: Optional[ModelMtpOutputCollector] = None # 用于判断 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 - # prompt_logics 用于在开启 return_all_prompt_logics 模式(如 enable_prompt_logprobs)时, # 保存整个 prefill 阶段每一个 token 位置对应的 logits(而非仅最后一个位置的 logits)。 # 此时 logits 依然只保存每个请求最后一个位置的 logits,prompt_logics 为可选项,仅在 # 需要返回 prompt logprobs 信息时才会非空。 prompt_logics: Optional[torch.Tensor] = None + def __post_init__(self) -> None: + 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) - if self.mtp_main_output_hiddens is not None: - self.mtp_main_output_hiddens = tensor_to_no_ref_tensor(self.mtp_main_output_hiddens) + self.mtp_collector.to_no_ref_tensor() diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 949384b437..5849cccf54 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -1,5 +1,6 @@ import os import torch +import torch.distributed as dist import copy import bisect import triton @@ -22,49 +23,61 @@ 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): + def gen_cuda_graph_batch_sizes( + 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, + ): args = get_env_start_args() - mtp_size = args.mtp_step + 1 - - # 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 + # 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. - 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): - batch_sizes.append(_batch_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}) - batch_sizes = list(set([e for e in batch_sizes if e < max_batch_size])) - batch_sizes.append(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] - batch_sizes = list(set(batch_sizes)) - batch_sizes.sort() - + 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): + 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, + 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.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.capture_infer_cost = capture_infer_cost + self.infer_cost_ms_by_batch_size = {} self.cuda_graph_batch_sizes = self.gen_cuda_graph_batch_sizes( - max_batch_size=max_batch_size, - tp_world_size=tp_world_size, + 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, ) - assert self.cuda_graph_batch_sizes[-1] == self.max_batch_size logger.info(f"cuda graph batch_sizes: {self.cuda_graph_batch_sizes}") def can_run(self, batch_size, max_len_in_batch): @@ -114,6 +127,7 @@ def _capture_decode(self, decode_func, infer_state: InferStateInfo): model_output = decode_func(infer_state) self.graph[batch_size] = (graph_obj, infer_state, model_output) graph_obj.replay() + self._measure_replay_cost(graph_obj=graph_obj, batch_size=batch_size) return model_output def _capture_decode_overlap( @@ -154,8 +168,32 @@ def _capture_decode_overlap( model_output1, ) graph_obj.replay() + 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) + 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( self, decode_func, @@ -211,7 +249,6 @@ 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 @@ -222,9 +259,10 @@ def warmup(self, model): 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_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, @@ -236,6 +274,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)], @@ -268,7 +308,6 @@ 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]: @@ -281,9 +320,10 @@ def warmup_overlap(self, model): 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_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, @@ -296,6 +336,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/hidden_collector.py b/lightllm/common/basemodel/hidden_collector.py new file mode 100644 index 0000000000..3eb946fe82 --- /dev/null +++ b/lightllm/common/basemodel/hidden_collector.py @@ -0,0 +1,238 @@ +from __future__ import annotations + +import copy +from abc import ABC, abstractmethod +from typing import List, Optional + +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 + + +class HiddenCollector(ABC): + """Hidden state 收集器的抽象基类。 + + 推理过程中,模型会在每一层计算完成后调用 :meth:`add`,并在一次 forward + 结束时调用 :meth:`finish_output` 生成统一的辅助输出。不同投机解码模式可以 + 通过子类决定不收集 hidden、只返回最终层 hidden、收集若干中间层 hidden, + 或同时收集 MTP head 生成的 token 与置信度。 + + 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_output` 后独立清理自己的容器。 + + 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_output` 前调用该接口。基类 + 默认不保存 tensor;需要直接返回最终层 hidden 的子类应重写该方法并保存 + 引用,随后在 :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_output(self, infer_state) -> ModelMtpOutputCollector: + """结束当前 microbatch 的收集并生成统一的 MTP 输出。 + + 子类应在此完成必要的拼接、TP/SP all-gather、DP unbalance 和 contiguous + 转换,将结果封装为 :class:`ModelMtpOutputCollector`,并在返回前清理当前 + 实例的临时状态。不需要提供额外 MTP 输出的实现应返回一个空 collector。 + + Args: + infer_state: 当前 forward 的推理状态,包含通信拓扑、DP balance 等信息。 + + Returns: + 本次 forward 的 MTP 辅助输出;未启用 MTP 时返回内容为空的 collector。 + """ + 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_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): + """Returns the final decoder hidden state without per-layer bookkeeping.""" + + 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_output(self, infer_state) -> ModelMtpOutputCollector: + assert self.final_hidden is not None + final_hidden = self.final_hidden + self.final_hidden = None + return ModelMtpOutputCollector(spec_hidden=final_hidden.contiguous()) + + +class LayerHiddenCollector(HiddenCollector): + """Collects selected decoder-layer outputs for an intermediate-hidden draft.""" + + def __init__(self, model) -> None: + self.model = model + self.layer_num = model.layers_num + self.layer_ids = self._load_layer_ids() + self.layer_hiddens: List[torch.Tensor] = [] + + 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: + 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( + 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 + # 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 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 ModelMtpOutputCollector(spec_hidden=hidden.contiguous()) diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index 10c35759aa..91c6e99699 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -34,19 +34,16 @@ 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 - # 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 # 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 @@ -68,6 +65,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/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/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py new file mode 100644 index 0000000000..be6c477b99 --- /dev/null +++ b/lightllm/common/basemodel/mtp_manager.py @@ -0,0 +1,98 @@ +from typing import ClassVar, Optional + +from lightllm.common.basemodel.hidden_collector import ( + FinalHiddenCollector, + HiddenCollector, + LayerHiddenCollector, + MtpHeadOutputCollector, + NoopHiddenCollector, +) +from lightllm.utils.envs_utils import get_env_start_args + + +class MtpManager: + """Manage MTP layout policy and model-local helper construction.""" + + _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 1 + + # 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 + + 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) + + 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: + 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) + 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..bf6039a48f 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__() @@ -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 @@ -132,13 +147,20 @@ 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] + 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 + ] 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_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) graph_infer_state.prefill_replay(infer_state) return graph_output_tensors @@ -177,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") @@ -191,11 +214,11 @@ 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, 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), ) @@ -237,6 +260,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") @@ -251,11 +275,11 @@ 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, 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/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/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/build_chained_mtp_decode_input.py b/lightllm/common/basemodel/triton_kernel/build_chained_mtp_decode_input.py new file mode 100644 index 0000000000..80da12fa5c --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/build_chained_mtp_decode_input.py @@ -0,0 +1,74 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _build_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 build_chained_mtp_decode_input_inplace( + 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 + + _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, + accept_len=accept_len, + BLOCK_SIZE=16, + num_warps=1, + num_stages=1, + ) + return draft_token_ids 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 new file mode 100644 index 0000000000..c93939fb8b --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/dynamic_mtp_utils.py @@ -0,0 +1,352 @@ +"""仅动态 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 + + +# 动态 verify 行选择。 +@triton.jit +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_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) + + scores = tl.load(base_ptr + offset, mask=store_mask, other=0.0) + # offset 0 是 target sample,本轮恒接受;只有 draft 条件接受概率需要 clamp。 + 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. + scores = tl.where((offset != 0) & (scores >= 0.99), 0.99, scores) + scores = tl.where((offset != 0) & (scores <= 0.01), 0.01, scores) + + cumulative_scores = tl.cumprod(scores, axis=0) + + tl.store(base_ptr + offset, cumulative_scores, mask=store_mask) + return + + +def sample_dynamic_mtp_row_mask( + dynamic_batch_size: int, + b_req_idx: torch.Tensor, + req_to_next_token_scores: torch.Tensor, + max_draft_step: int, + pre_draft_step: int = None, +) -> torch.Tensor: + dynamic_batch_size = int(dynamic_batch_size) + 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_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 + + # 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), + num_warps=1, + num_stages=1, + ) + + 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 + + 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, + 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, + 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) + + 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 + + +@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 _compact_decode_model_input( + model_input: ModelInput, + selected_row_mask: torch.Tensor, + dynamic_batch_size: 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 + 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) + 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 = 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 + ) + 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, + 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, + 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_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 + 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, + ) + # 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/fa3_utils.py b/lightllm/common/basemodel/triton_kernel/fa3_utils.py index 0a524b63b6..3d04558273 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_SPEC_FA3_FAST_PATH_MAX_BATCH_SIZE = 1024 +_DYNAMIC_SPEC_FA3_COMPACT_BLOCK_SIZE = 256 + + @triton.jit def page_table_copy_kernel( page_table_ptr, @@ -57,6 +62,214 @@ def page_table_copy( ) +@triton.jit +def _build_dynamic_spec_fa3_decode_params_kernel( + b_req_idx, + b_seq_len, + b_mark_mtp_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 + + 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 + + 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) + + 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) + + +@triton.jit +def _count_dynamic_spec_fa3_decode_params_kernel( + b_mark_mtp_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 + + 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) + + +@triton.jit +def _compact_dynamic_spec_fa3_decode_params_kernel( + b_req_idx, + b_seq_len, + b_mark_mtp_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_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 + 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_spec_fa3_decode_params( + b_req_idx: torch.Tensor, + b_seq_len: 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_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, 2, 2] + b_mark_mtp_shared_group = [ 0, 0, 3, 1, 0, 2, 1, 1] + + 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``:: + + # 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 + 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_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 + + 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_spec_fa3_decode_params_kernel[(1,)]( + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + 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, + 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_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_mtp_shared_group.device) + + _count_dynamic_spec_fa3_decode_params_kernel[grid]( + b_mark_mtp_shared_group=b_mark_mtp_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) + + _compact_dynamic_spec_fa3_decode_params_kernel[grid]( + b_req_idx=b_req_idx, + b_seq_len=b_seq_len, + 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, + 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 + + def test_page_table_copy(): import torch 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/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/mtp_state_params.py b/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py new file mode 100644 index 0000000000..b81bf9ec16 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/linear_att/mtp_state_params.py @@ -0,0 +1,111 @@ +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) + + 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 + # 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() + + # 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( + 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]: + """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 + 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 + 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/basemodel/triton_kernel/mtp_utils.py b/lightllm/common/basemodel/triton_kernel/mtp_utils.py index 26e1468bd4..7e943f2925 100644 --- a/lightllm/common/basemodel/triton_kernel/mtp_utils.py +++ b/lightllm/common/basemodel/triton_kernel/mtp_utils.py @@ -1,7 +1,99 @@ +"""固定布局与动态布局共同使用的 MTP Triton 算子。""" + +from typing import Optional + import triton 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( @@ -12,14 +104,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) @@ -52,22 +148,23 @@ 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,) + 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. """ - 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] + 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 @@ -79,7 +176,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, @@ -91,12 +188,19 @@ 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, + schedule_scores_stride, mtp_accept_len, b_req_mtp_start_loc, b_req_idx, - mtp_step, + draft_step, + verify_width, + HAS_NEXT_TOKEN_SCORES: tl.constexpr, BLOCK_SIZE: tl.constexpr, ): @@ -105,43 +209,100 @@ def _fwd_kernel_mtp_scatter_next_token_ids( 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 + + # 第 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, + ) - 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, + # 从第 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, - scatter_next_token_ids, - mask=offset < mtp_step, + 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: + # 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 + cur_index * schedule_scores_stride + offset, + mask=offset < draft_step, + other=0.0, + ) + tl.store( + req_to_next_token_scores + cur_req_idx * req_to_next_token_scores_stride + offset + 1, + draft_scores, + mask=offset + 1 < 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_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] ): - max_mtp_step = req_to_next_token_ids.shape[1] + """将 target token 和 draft proposal 写入请求级 MTP buffer。""" + + 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] + 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 == 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. + 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 + ) + 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 _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, + schedule_scores_stride=schedule_scores_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, + draft_step=draft_step, + verify_width=verify_width, + HAS_NEXT_TOKEN_SCORES=HAS_NEXT_TOKEN_SCORES, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps, num_stages=1, @@ -217,7 +378,7 @@ def linear_att_mtp_state_index_update( 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. @@ -227,10 +388,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] @@ -258,14 +419,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/common/basemodel/triton_kernel/select_mtp_rows.py b/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py new file mode 100644 index 0000000000..5e230ce855 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/select_mtp_rows.py @@ -0,0 +1,187 @@ +"""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 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,)), + ) + if req_num == 0: + return selected + + 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/common/kv_cache_mem_manager/operator/linear_att.py b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py index 147b43f697..49e2549265 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py +++ b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py @@ -191,7 +191,7 @@ def offload_gpu_kv_to_cpu_cache( 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 + 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 c8f69dac84..907cc494a6 100644 --- a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py @@ -29,11 +29,7 @@ def __init__( super().__init__(size, dtype, num_kv_heads, head_dim, full_att_layer_num, always_copy, mem_fraction) 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): 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/req_manager.py b/lightllm/common/req_manager.py index 070da7412f..3de7de8f12 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -113,14 +113,21 @@ 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", ) + 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 + ) + self.req_to_exponential_decay_length_penalty = torch.zeros( max_request_num + 1, dtype=torch.float32, device="cuda" ) @@ -137,6 +144,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_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/__init__.py b/lightllm/models/__init__.py index f619b1d88f..c7e9a59aad 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -43,4 +43,16 @@ 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 .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 e2b2a56137..6668678e6a 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.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 +@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/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 2e4ba5c86b..06f941d806 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.draft_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..6a3c32ebbc 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.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 @@ -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/qwen2_vl/infer_struct.py b/lightllm/models/qwen2_vl/infer_struct.py index 04f7bc3895..0c61655091 100644 --- a/lightllm/models/qwen2_vl/infer_struct.py +++ b/lightllm/models/qwen2_vl/infer_struct.py @@ -17,18 +17,34 @@ 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) + + # 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) + + # 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. else: - b_position_delta = self.b_position_delta.to(dtype=self.position_ids.dtype) - 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, "decode requires b_position_delta" + 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/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..3d2c613c8e --- /dev/null +++ b/lightllm/models/qwen3_5_dflash/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,39 @@ +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): + 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_, + ) diff --git a/lightllm/models/qwen3_5_dflash/model.py b/lightllm/models/qwen3_5_dflash/model.py new file mode 100644 index 0000000000..ae376ea334 --- /dev/null +++ b/lightllm/models/qwen3_5_dflash/model.py @@ -0,0 +1,49 @@ +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.qwen3_dflash.model import Qwen3DFlashModel +from lightllm.models.draft_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): + """Adapter for a Qwen3 DFlash checkpoint paired with a Qwen3.5 target.""" + + pre_and_post_weight_class = Qwen35DFlashPreAndPostLayerWeight + + def _init_config(self): + super()._init_config() + self.config.update(self.config.get("dflash_config", {})) + + 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 parallel block drafter requires matching draft and target KV shapes, " + f"got draft={draft_kv_shape}, target={target_kv_shape}." + ) + 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_ 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..865527597c --- /dev/null +++ b/lightllm/models/qwen3_5_dspark/model.py @@ -0,0 +1,41 @@ +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.models.qwen3_dspark.model import Qwen3DSparkModel +from lightllm.models.draft_registry import DraftModelRegistry + + +@DraftModelRegistry(model_type=("qwen3_5", "qwen3_5_text"), spec_modes="dspark") +class Qwen3_5DSparkModel(Qwen3DSparkModel): + """Adapter for the current Qwen3 DSpark checkpoint with a Qwen3.5 target.""" + + def _init_config(self): + super()._init_config() + self.config.update(self.config.get("dflash_config", {})) + + 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"] + + # 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. + 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 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_moe_mtp/model.py b/lightllm/models/qwen3_5_moe_mtp/model.py index 022864f6b3..e852e24c59 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.draft_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..b8639a9970 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.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 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/__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..b33aac3179 --- /dev/null +++ b/lightllm/models/qwen3_dflash/infer_struct.py @@ -0,0 +1,5 @@ +from lightllm.models.llama.infer_struct import LlamaInferStateInfo + + +class Qwen3DFlashInferStateInfo(LlamaInferStateInfo): + """DFlash attention metadata.""" 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..e69de29bb2 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..b140c80d5f --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_infer/post_layer_infer.py @@ -0,0 +1,15 @@ +import torch + +from lightllm.models.llama.layer_infer.post_layer_infer import LlamaPostLayerInfer + + +class Qwen3DFlashPostLayerInfer(LlamaPostLayerInfer): + def token_forward(self, input_embdings: torch.Tensor, infer_state, layer_weight): + 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, + 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..bafe533a91 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_infer/pre_layer_infer.py @@ -0,0 +1,26 @@ +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): + """Project target hiddens for DFlash commit prefill.""" + + def __init__(self, network_config): + super().__init__(network_config) + self.eps_ = network_config["rms_norm_eps"] + + def context_forward( + self, + input_ids, + infer_state, + layer_weight: Qwen3DFlashPreAndPostLayerWeight, + ): + 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 new file mode 100644 index 0000000000..84ee0b4059 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_infer/transformer_layer_infer.py @@ -0,0 +1,69 @@ +import torch + +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 + + +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.partial_rotary_factor = network_config.get("partial_rotary_factor", 1.0) + + def context_forward( + self, + input_embdings: torch.Tensor, + infer_state: Qwen3DFlashInferStateInfo, + layer_weight: Qwen3DFlashTransformerLayerWeight, + ) -> torch.Tensor: + token_num, _ = input_embdings.shape + 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( + cache_kv[:, : self.tp_k_head_num_, :], + 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 + + def _get_qkv(self, input, infer_state: Qwen3DFlashInferStateInfo, layer_weight: Qwen3DFlashTransformerLayerWeight): + 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) + + 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, + partial_rotary_factor=self.partial_rotary_factor, + ) + 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..ce2ca006e1 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,62 @@ +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_, + ) 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..0eaf6b8684 --- /dev/null +++ b/lightllm/models/qwen3_dflash/layer_weights/transformer_layer_weight.py @@ -0,0 +1,98 @@ +from lightllm.common.basemodel.layer_weights.meta_weights import ( + COLMMWeight, + KVROWNMMWeight, + QKRMSNORMWeight, + 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" + + 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"), + ) + + 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() + 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_dflash/model.py b/lightllm/models/qwen3_dflash/model.py new file mode 100644 index 0000000000..122f4044c3 --- /dev/null +++ b/lightllm/models/qwen3_dflash/model.py @@ -0,0 +1,121 @@ +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 +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 + + +@DraftModelRegistry(model_type="qwen3", spec_modes="dflash") +class Qwen3DFlashModel(LlamaTpPartModel): + """Qwen3 DFlash draft model.""" + + is_mtp_draft_model = True + + 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) + + 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") + + 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 + self._sin_cached = self.main_model._sin_cached + self.block_size = int(self.config["block_size"]) + self.mask_token_id = int(self.config["mask_token_id"]) + + def _init_req_manager(self): + self.req_manager = self.main_model.req_manager + + def _init_mem_manager(self): + # Draft KV uses target-owned cache slots. + self.mem_manager = self.main_model.mem_manager + + def _init_att_backend(self): + super()._init_att_backend() + # 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") + + def _init_infer_layer(self, start_layer_index=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) + + def _init_weights(self, start_layer_index=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"]) + ] + + 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} + + def _autotune_warmup(self): + return + + def _init_padded_req(self): + return + + def _init_prefill_cuda_graph(self): + self.prefill_graph = None 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..e69de29bb2 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..5a74cd988e --- /dev/null +++ b/lightllm/models/qwen3_dspark/layer_infer/post_layer_infer.py @@ -0,0 +1,195 @@ +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.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_ = 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: + return torch.nn.functional.embedding(token_ids, layer_weight.markov_w1_weight_.weight) + + def _markov_step_latent( + self, + prev_embeddings: torch.Tensor, + hidden_states: torch.Tensor, + state: torch.Tensor, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ): + if self.markov_head_type_ == "vanilla": + 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(layer_weight.markov_gate_proj_weight_.mm(gate_input)) + return state, gate * prev_embeddings + + if state is None: + state = torch.zeros_like(prev_embeddings) + joint_input = torch.cat([state, prev_embeddings, hidden_states], dim=-1) + 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, torch.tanh(output_raw) + + @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.enable_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 = layer_weight.confidence_head_weight_.mm(features.flatten(0, -2)) + return logits.float().view(features.shape[:-1]) + + def _sample_markov( + self, + local_logits: torch.Tensor, + block_hidden: torch.Tensor, + infer_state: Qwen3DFlashInferStateInfo, + anchor_token_ids: torch.Tensor, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ) -> torch.Tensor: + """Run sequential Markov decoding over TP-local vocabulary logits.""" + + 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) + 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) + 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, + infer_state: Qwen3DFlashInferStateInfo, + layer_weight: Qwen3DSparkPreAndPostLayerWeight, + ): + if infer_state.is_prefill: + return super().token_forward( + input_embdings=input_embdings, + infer_state=infer_state, + layer_weight=layer_weight, + ) + + 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, + ) + 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) + 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/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..437cd36827 --- /dev/null +++ b/lightllm/models/qwen3_dspark/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,83 @@ +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 + + +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: + # 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_, + tp_rank=0, + tp_world_size=1, + ) + self.markov_w2_weight_ = LMHeadWeight( + dim=markov_rank, + vocab_size=vocab_size, + weight_name="markov_head.markov_w2.weight", + data_type=self.data_type_, + ) + if self.markov_head_type == "gated": + 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_, + # W8A8 MM does not support the Markov projection bias. + quant_method=None, + tp_rank=0, + tp_world_size=1, + ) + elif self.markov_head_type == "rnn": + 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_, + quant_method=None, + tp_rank=0, + tp_world_size=1, + ) + 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_ = 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_, + # 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/models/qwen3_dspark/model.py b/lightllm/models/qwen3_dspark/model.py new file mode 100644 index 0000000000..a9e49ee91f --- /dev/null +++ b/lightllm/models/qwen3_dspark/model.py @@ -0,0 +1,17 @@ +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 + + +@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 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..e69de29bb2 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..80290cfae4 --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_infer/pre_layer_infer.py @@ -0,0 +1,38 @@ +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 and fixed-width draft hidden preparation.""" + + def __init__(self, network_config): + super().__init__(network_config) + self.hidden_size_ = network_config["hidden_size"] + + def prepare_spec_draft_hiddens( + self, + infer_state: InferStateInfo, + ) -> None: + 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, + input_ids, + infer_state: InferStateInfo, + layer_weight: Qwen3EaglePreAndPostLayerWeight, + ): + self.prepare_spec_draft_hiddens(infer_state) + return super().context_forward(input_ids, infer_state, layer_weight) + + def token_forward( + self, + input_ids, + infer_state: InferStateInfo, + layer_weight: Qwen3EaglePreAndPostLayerWeight, + ): + self.prepare_spec_draft_hiddens(infer_state) + 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..4f34a9e76b --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_infer/transformer_layer_infer.py @@ -0,0 +1,54 @@ +import torch + +from lightllm.common.basemodel.infer_struct import InferStateInfo +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(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_( + input=infer_state.eagle_draft_hidden_states, + eps=self.eps_, + alloc_func=self.alloc_tensor, + ) + 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_), + cache_kv[:, : self.tp_k_head_num_, :], + infer_state.position_cos, + infer_state.position_sin, + ) + 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..12ef341d51 --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,67 @@ +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_, + ) + + # 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: + 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_, + ) 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..53b6d5cc0d --- /dev/null +++ b/lightllm/models/qwen3_eagle/layer_weights/transformer_layer_weight.py @@ -0,0 +1,61 @@ +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 + + +class Qwen3EagleTransformerLayerWeight(LlamaTransformerLayerWeight): + def _init_weight_names(self): + super()._init_weight_names() + 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" + 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_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_ = 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..8ab63705aa --- /dev/null +++ b/lightllm/models/qwen3_eagle/model.py @@ -0,0 +1,117 @@ +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 +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 + + 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) + + 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 _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 + + def _init_req_manager(self): + self.req_manager = self.main_model.req_manager + + def _init_mem_manager(self): + self.mem_manager = self.main_model.mem_manager + + def _init_weights(self, start_layer_index=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"]) + ] + # 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): + 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) + + 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: + 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..88522a5e23 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.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 @@ -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..1b513cc27d 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, Type, TypeVar, Union + from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index edcf04f710..ffc5434e28 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -737,31 +737,36 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "eagle_with_att", "vanilla_no_att", "eagle_no_att", + "eagle3", + "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 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", 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", 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="""Enable dynamic speculative scheduling.""", ) parser.add_argument( "--kv_quant_calibration_config_path", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 6cdbfe0f9e..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) @@ -164,6 +167,11 @@ def _launch_subprocesses(args: StartArgs): # mtp params check if args.mtp_mode is not None: if args.mtp_draft_model_dir is None: + 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: diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index b268120c90..39177b1498 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -133,6 +133,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 +205,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 ecc447d6aa..ee0c9603b7 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -188,12 +188,16 @@ class StartArgs: "eagle_with_att", "vanilla_no_att", "eagle_no_att", + "eagle3", + "dspark", + "dflash", 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 195ef17e45..1b2300694b 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -711,6 +711,8 @@ 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] = {} first_token_cost_ms = sys.float_info.max prompt_tokens = len(prompt_ids) is_first_token = True @@ -744,6 +746,10 @@ 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) + 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_first_token: first_token_cost_ms = (time.time() - start_time) * 1000 @@ -773,9 +779,14 @@ 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_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} " @@ -793,7 +804,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_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_accepted_tokens_per_step:{mtp_avg_accepted_tokens_per_step} " + f"mtp_avg_verify_tokens_per_step:{mtp_avg_verify_tokens_per_step} " ) self.metric_client.histogram_observe("lightllm_cache_length", prompt_cache_len) @@ -938,6 +954,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 96d1361e63..5912075cea 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -475,6 +475,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 @@ -488,6 +489,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 @@ -506,9 +508,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_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/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/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index e0a7ebae77..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 @@ -859,6 +860,12 @@ 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): + self.shm_req.mtp_verify_token_num += verify_token_num + + 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): 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/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index eaf3552607..e0d1b3924a 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -1,4 +1,5 @@ import os + import numpy as np import torch import time @@ -8,7 +9,7 @@ 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 @@ -19,7 +20,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 @@ -43,10 +43,6 @@ ) 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 @@ -71,6 +67,7 @@ def __init__(self) -> 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 @@ -245,9 +242,9 @@ 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: - self.init_mtp_draft_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) @@ -301,23 +298,22 @@ 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 + 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") 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}" + 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(num_mtp_modules): - mtp_model_cfg, _ = PretrainedConfig.get_config_dict(self.args.mtp_draft_model_dir[i]) - model_type = mtp_model_cfg.get("model_type", "") - mtp_model_kvargs = { - "weight_dir": self.args.mtp_draft_model_dir[i], + for i in range(draft_model_count): + draft_model_cfg, _ = PretrainedConfig.get_config_dict(draft_model_dirs[i]) + 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), @@ -326,7 +322,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": main_kvargs.get("graph_max_batch_size", 16), + "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"], @@ -339,33 +335,13 @@ def init_mtp_draft_model(self, main_kvargs: dict): "mtp_previous_draft_models": self.draft_models.copy(), } - 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"] - 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"] - self.draft_models.append(Qwen3MOEMTPModel(mtp_model_kvargs)) - elif model_type == "mistral": - assert self.args.mtp_mode in ["vanilla_no_att", "eagle_no_att"] - 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"] - 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"] - 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"] - from lightllm.models.qwen3_5_moe_mtp.model import Qwen3_5MoeMTPModel - - self.draft_models.append(Qwen3_5MoeMTPModel(mtp_model_kvargs)) - else: - raise ValueError(f"Unsupported MTP model type: {model_type}") + draft_model_class = get_draft_model_class( + model_cfg=draft_model_cfg, + spec_mode=spec_mode, + ) + self.draft_models.append(draft_model_class(draft_model_kvargs)) - self.logger.info(f"loaded mtp model class {self.draft_models[i].__class__}") + self.logger.info(f"loaded speculative draft model class {self.draft_models[i].__class__}") return def _async_copy_next_token_infos_to_pin_mem( @@ -817,9 +793,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, @@ -828,6 +804,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 ): @@ -858,31 +838,15 @@ 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) - return - def _gen_argmax_token_ids(self, model_output: ModelOutput): logits = model_output.logits - draft_next_token_ids_gpu = torch.argmax(logits, dim=-1) - 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) + return draft_next_token_ids_gpu, max_probs def _sample_and_scatter_token( self, 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..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 @@ -1,7 +1,7 @@ import torch import time -from typing import List, Optional, Callable, Dict, Any -from queue import Queue +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 @@ -9,22 +9,16 @@ 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.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 lightllm.utils.envs_utils import get_env_start_args from .control_state import ControlState +from lightllm.utils.envs_utils import get_env_start_args logger = init_logger(__name__) @@ -37,12 +31,9 @@ def __init__(self) -> None: self.control_state_machine = ControlState() # 在 mtp 模式下切换绑定的prefill 和 decode 函数 - if get_env_start_args().mtp_mode: + if get_env_start_args().mtp_mode is not None: 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 else: self.prefill = self.prefill_normal self.decode = self.decode_normal @@ -50,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: @@ -209,8 +208,11 @@ 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_engine = self.spec_engine + spec_engine.fill_draft_model_kv_state( + target_model_input=model_input, + target_model_output=model_output, + 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, @@ -246,38 +248,43 @@ def decode_mtp( event_pack: OverlapEventPack, decode_reqs: List[InferReq], ): - """ - MTP解码的通用流程,整合eagle和vanilla的共同逻辑 - """ + """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()): - b_mtp_index_cpu = model_input.b_mtp_index + 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) - next_token_ids, next_token_logprobs = sample(model_output.logits, run_reqs, self.eos_id) + # 动态 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() + 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, + 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_mtp_start_loc = gen_b_req_mtp_start_loc(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=model_input.b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, + b_mtp_index=model_input.b_mtp_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, - 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, @@ -286,49 +293,80 @@ def decode_mtp( key="mtp_accept_len", gpu_tensor=mtp_accept_len, ) + 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( + 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, + ) + 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_logprobs, next_token_ranks) - - # 调用具体的draft decode函数 - additional_mem_indexes_cpu = self._draft_decode_func( - main_model_input=model_input, - main_model_output=model_output, + ) = self._async_copy_next_token_infos_to_pin_mem( next_token_ids=next_token_ids, - mtp_accept_len=mtp_accept_len, - b_req_mtp_start_loc=b_req_mtp_start_loc, + next_token_logprobs=next_token_logprobs, + 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() # 第二阶段 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] + + # 当 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) # 第三阶段 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_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, + ) - 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") + 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], @@ -338,109 +376,17 @@ def decode_mtp( 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) + proposal.extra_mem_indexes_cpu.append( + MtpMemIndexesToFree( + mem_indexes_cpu=model_input.mem_indexes_cpu, + free_mask_cpu=accepted_index_cpu == 0, + ), + ) + mtp_utils.free_mem_indexes( + backend=self, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, + ) # 第四阶段 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..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 @@ -21,12 +21,13 @@ class DiversehBackend(ChunkedPrefillBackend): def __init__(self) -> None: super().__init__() - if get_env_start_args().mtp_mode: - # 当前只有 mistral mtp 可以使用 diverse mode 的 mtp 功能。 - self.prefill = self.beam_prefill - assert get_env_start_args().mtp_mode in ["vanilla_no_att", "eagle_no_att"] - else: - self.prefill = self.beam_prefill + self.prefill = self.beam_prefill + spec_mode = get_env_start_args().mtp_mode + if spec_mode is not None: + assert spec_mode in [ + "vanilla_no_att", + "eagle_no_att", + ] 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 f6ca89e651..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 @@ -1,29 +1,25 @@ 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.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 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, - padded_prepare_decode_inputs, - padded_overlap_prepare_prefill_inputs, - padded_overlap_prepare_decode_inputs, + prepare_prefill_inputs, + 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.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 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 from .control_state import DPControlState @@ -35,21 +31,18 @@ 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 + 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 parallel block drafting yet.") 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.is_mtp_eagle 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 else: if self.enable_prefill_microbatch_overlap: self.prefill = self.prefill_overlap @@ -64,6 +57,31 @@ 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, + ) + # 非 overlap DP 与普通后端复用同一个 SpecEngine。 + self.spec_engine = SpecEngine( + backend=self, + spec_mode=self.args.mtp_mode, + enable_dynmaic_mtp=self.args.mtp_dynamic_verify, + ) + + 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 + ) + self.decode_draft_engine = ( + self.dp_overlap_spec_engine if self.enable_decode_microbatch_overlap else self.spec_engine + ) + return + def _init_reqs(self, reqs: List[Tuple]): if not self.args.enable_dp_prompt_cache_fetch: return super()._init_reqs(reqs) @@ -157,7 +175,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) @@ -169,16 +187,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() @@ -210,7 +228,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()): @@ -222,9 +240,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, @@ -261,11 +279,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) @@ -276,18 +292,15 @@ 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: + if req_num0 + req_num1 > 0: ( _, next_token_ids_cpu, @@ -309,7 +322,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) @@ -336,32 +349,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, @@ -378,7 +375,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) @@ -405,14 +402,17 @@ 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) - 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: ( @@ -421,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, @@ -429,15 +429,15 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] b_prefill_has_output_cpu=b_has_out_cpu, mask_func=None, ) - - # 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( - model_input=model_input, - model_output=model_output, - next_token_ids=draft_next_token_ids_gpu, + else: + next_token_ids = torch.empty((0,), dtype=torch.int64, device=model_input.b_req_idx.device) + + # 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=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) @@ -474,48 +474,49 @@ 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) - b_mtp_index_cpu = model_input.b_mtp_index - req_num = len(run_reqs) + """复用普通 SpecEngine 执行 DP speculative draft-and-verify。""" + + model_input, run_reqs = prepare_decode_inputs(req_objs=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, + 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) 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) + 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) - mtp_accept_len, accepted_index = self._verify_mtp_v2( - new_next_token_ids=next_token_ids, - b_req_idx=b_req_idx, + 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=model_input.b_req_idx, b_req_mtp_start_loc=b_req_mtp_start_loc, + b_mtp_index=model_input.b_mtp_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, - 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, - ) accepted_index_cpu = g_pin_mem_manager.async_copy_from_gpu_tensor( key="accepted_index", gpu_tensor=accepted_index, @@ -524,23 +525,52 @@ 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, + ) + 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() - eagle_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() @@ -549,19 +579,39 @@ 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() - verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] + 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() - 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") + 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, + ) + + proposal.extra_mem_indexes_cpu.append( + MtpMemIndexesToFree( + mem_indexes_cpu=model_input.mem_indexes_cpu, + free_mask_cpu=accepted_index_cpu == 0, + ), + ) + + 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], @@ -570,129 +620,31 @@ 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, + 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() + sync_event.synchronize() + mtp_utils.free_mem_indexes( + backend=self, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, + ) 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, - ): - 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) - - 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 - # 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, - 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, - ) - 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, - req_num: int, - ): - 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) - - 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 = 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 - # 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 = 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, - 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, - ) - return eagle_mem_indexes_cpu - def prefill_overlap_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq]): ( 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) @@ -700,18 +652,21 @@ 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) - 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) + 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, 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: + if req_num > 0: ( next_token_ids, next_token_ids_cpu, @@ -725,49 +680,28 @@ 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 = 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, + target_model_output0=model_output0, + target_next_token_ids0=target_next_token_ids_gpu0, + target_model_input1=model_input1, + target_model_output1=model_output1, + target_next_token_ids1=target_next_token_ids_gpu1, + ) - # 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_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) - - 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) - g_infer_context.copy_linear_att_state_to_cache_buffer(b_req_idx=_b_req_idx, reqs=run_reqs) + if req_num > 0 and g_infer_context.is_linear_att_mixed_model: + 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) @@ -794,28 +728,59 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf ( model_input0, run_reqs0, - _, + decode_reqs0, model_input1, run_reqs1, - _, - ) = 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 + decode_reqs1, + ) = overlap_prepare_decode_inputs(req_objs=decode_reqs) + 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()): - + 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 - b_req_idx, mtp_accept_len, b_req_mtp_start_loc, next_token_ids = None, None, None, None - if (req_num0 + req_num1) > 0: + if req_num > 0: + assert len(run_reqs) == verify_row_num logits = torch.empty( - (req_num0 + req_num1, logits0.shape[1]), dtype=logits0.dtype, device=logits0.device + (verify_row_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) + 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) ( @@ -824,32 +789,24 @@ 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) 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 = 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, ) - 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_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, @@ -858,25 +815,44 @@ 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) - + 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() verify_event.record() - 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, - b_req_mtp_start_loc=b_req_mtp_start_loc, - req_num0=req_num0, - req_num1=req_num1, + 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, + accept_len0=mtp_accept_len0, + target_model_input1=model_input1, + target_model_output1=model_output1, + target_next_token_ids1=target_next_token_ids1, + accept_len1=mtp_accept_len1, + draft_step=spec_plan.draft_step, ) + 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, @@ -885,23 +861,48 @@ 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() - verify_ok_reqs = [run_reqs[i] for i in range(len(run_reqs)) if accepted_index_cpu[i] == 1] + mtp_utils.record_request_mtp_metrics( + backend=self, + 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_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() 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, + ) + 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, + ), + ) ) - 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) - 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") + 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], @@ -910,182 +911,18 @@ 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, + 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_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, + sync_event.synchronize() + mtp_utils.free_mem_indexes( + backend=self, + extra_mem_indexes_cpu=proposal.extra_mem_indexes_cpu, ) - 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) + 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, - ): - 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_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 - ) - - # 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_output0, draft_model_output1 = self.draft_models[draft_model_idx].microbatch_overlap_decode( - 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) - 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) - - 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, - 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 None - - 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, - ): - 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_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 - ) - 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 = 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.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_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 - 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 = 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 = 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) - 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) - - 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, - 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 eagle_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 deleted file mode 100644 index 68af30b505..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/generic_padded_pre_process.py +++ /dev/null @@ -1,257 +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 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 - prefix_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 - prefix_total_token_num += req.cur_kv_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 - prefix_total_token_num += 0 - 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, - prefix_total_token_num=prefix_total_token_num, - 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) - - # 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) - - if padded_mem_indexes_num > 0: - mem_indexes = F.pad( - input=mem_indexes, - pad=(0, padded_mem_indexes_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, - 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, - 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 ae294544ce..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 @@ -3,16 +3,13 @@ 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]]: run_reqs = [] total_token_num = 0 - prefix_total_token_num = 0 input_ids = [] b_req_idx = [] b_seq_len = [] @@ -44,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"): @@ -53,11 +49,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") @@ -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, ) @@ -124,19 +121,24 @@ 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") 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,13 +157,63 @@ 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, ) 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 + decode_reqs0 = req_objs[:split_req_bound] + decode_reqs1 = req_objs[split_req_bound:] + model_input0, run_reqs0 = prepare_decode_inputs( + req_objs=decode_reqs0, + ) + model_input1, run_reqs1 = prepare_decode_inputs( + req_objs=decode_reqs1, + ) + return model_input0, run_reqs0, decode_reqs0, model_input1, run_reqs1, decode_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: @@ -172,48 +224,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/mode_backend/mtp_pre_process.py b/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py deleted file mode 100644 index 3ef0395431..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/mtp_pre_process.py +++ /dev/null @@ -1,24 +0,0 @@ -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( - 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 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/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/__init__.py new file mode 100644 index 0000000000..66dc8ab2fc --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/__init__.py @@ -0,0 +1,5 @@ +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__ = ["DPOverlapSpecEngine", "SpecEngine"] 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..5200db6e16 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -0,0 +1,186 @@ +from __future__ import annotations + +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_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.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 + + +class DPOverlapSpecEngine: + """双 microbatch overlap draft 流程使用的 DP MTP engine。""" + + 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, + ) + + 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, + 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: + self.proposer.fill_draft_model_kv_state_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_ids1=target_next_token_ids1, + ) + + 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] + 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] + 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_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, + accept_len1=accept_len1, + 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, + ) + + +__all__ = ["DPOverlapSpecEngine"] 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..8c0ff84b10 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/__init__.py @@ -0,0 +1,48 @@ +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: "ModeBackend", + 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..95286c7b70 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/base.py @@ -0,0 +1,51 @@ +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 ( + SpecProposal, +) + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + + +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( + 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: + """Build draft state from two overlapped target-prefill microbatches.""" + + raise NotImplementedError + + @abstractmethod + 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] + 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] + accept_len1: torch.Tensor, # [real_req_num1] + 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..6594182db7 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle3.py @@ -0,0 +1,21 @@ +import torch + +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(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 _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 _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 new file mode 100644 index 0000000000..077d478fdc --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_no_att.py @@ -0,0 +1,150 @@ +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.utils import ( + get_dp_overlap_req_start_rows, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + EagleSpecProposal, +) + + +class DpOverlapEagleNoAttProposer(BaseDpOverlapProposer): + """DP ``eagle_no_att`` proposer。""" + + 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: + pass + + def propose_next_overlap( + self, + 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, + accept_len1: torch.Tensor, + draft_step: int, + ) -> EagleSpecProposal: + 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) + + 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,) + 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_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py new file mode 100644 index 0000000000..6b8c23e8fd --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -0,0 +1,258 @@ +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.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): + """DP ``eagle_with_att`` proposer。""" + + 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: + 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, + 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, + accept_len1: torch.Tensor, + draft_step: int, + ) -> EagleSpecProposal: + """提交两个 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 + + 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 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 + + 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 new file mode 100644 index 0000000000..4c3261aa26 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_no_att.py @@ -0,0 +1,149 @@ +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.utils import ( + get_dp_overlap_req_start_rows, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import ( + VanillaSpecProposal, +) + + +class DpOverlapVanillaNoAttProposer(BaseDpOverlapProposer): + """DP ``vanilla_no_att`` proposer。""" + + 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: + pass + + def propose_next_overlap( + self, + 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, + accept_len1: torch.Tensor, + draft_step: int, + ) -> VanillaSpecProposal: + 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) + + 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,) + 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=[], + 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 new file mode 100644 index 0000000000..4f3c117b34 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/vanilla_with_att.py @@ -0,0 +1,221 @@ +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.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, +) +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, +) +from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager + + +class DpOverlapVanillaWithAttProposer(BaseDpOverlapProposer): + """DP ``vanilla_with_att`` proposer。""" + + 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: + 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_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) + + def propose_next_overlap( + self, + 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, + accept_len1: torch.Tensor, + draft_step: int, + ) -> VanillaSpecProposal: + assert draft_step == self.backend.max_draft_step + assert len(self.backend.draft_models) == draft_step + + 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=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 = 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 = [ + 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)) + 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): + 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 + 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()) + + if step + 1 < draft_step: + for batch_index, model_input in enumerate(model_inputs): + 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], + accept_len=accept_len_by_batch[batch_index], + ) + + return VanillaSpecProposal( + token_ids=proposal_token_ids, + extra_mem_indexes_cpu=[], + schedule_scores=proposal_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_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/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py new file mode 100644 index 0000000000..9d59afd8a7 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -0,0 +1,157 @@ +from __future__ import annotations + +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.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager +from lightllm.server.router.model_infer.mtp_speculative.planner import ( + BaseMtpPlanner, + DSparkPlanner, + FixedSpecPlanner, + LightSpecPlanner, + 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 BaseSpecProposer, SpecProposal + +if TYPE_CHECKING: + from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + + +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``. + """ + + 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, + backend=backend, + enable_dynmaic_mtp=enable_dynmaic_mtp, + ) + self.planner: BaseMtpPlanner = self._build_mtp_planner( + spec_mode=spec_mode, + enable_dynmaic_mtp=enable_dynmaic_mtp, + ) + + # Prefill draft-state initialization. + + 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, + ) + + # Decode planning. + + def plan_decode(self, model_input: ModelInput, decode_reqs: List) -> SpecDecodePlan: + """Return the fixed or dynamic speculative plan for one decode iteration.""" + + return self.planner.plan( + decode_reqs=decode_reqs, + origin_batch_size=model_input.batch_size, + ) + + 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 planned batch is smaller.""" + + 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.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。 + # 动态 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, + req_num=req_num, + dynamic_batch_size=plan.dynamic_batch_size, + 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, + ) + 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 + + # Draft proposal generation. + + 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] + draft_step: int, + 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=draft_step, + accept_len=accept_len, + ) + + # Planner runtime statistics. + + def update_planner_statics( + self, + plan: SpecDecodePlan, + proposal: SpecProposal, + req_num: int, + accept_lengths_cpu: torch.Tensor, + ) -> None: + """Update the current planner with iteration-level runtime statistics.""" + + self.planner.update_statics( + plan=plan, + proposal=proposal, + req_num=req_num, + accept_lengths=accept_lengths_cpu, + ) + + 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 spec_mode == "dspark": + return DSparkPlanner(backend=self.backend) + return LightSpecPlanner(spec_mode=spec_mode, backend=self.backend) 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..6db1b911be --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/__init__.py @@ -0,0 +1,13 @@ +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", + "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..c1becbdcd2 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/base.py @@ -0,0 +1,142 @@ +from __future__ import annotations + +from abc import ABC, abstractmethod +from dataclasses import dataclass +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: + """Planner decision for one target decode iteration. + + ``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. + + 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 + """ + + 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 + # the previous iteration. Such an iteration does not represent one + # well-defined LightSpec runtime configuration. + all_reqs_have_proposals: bool = True + + @property + def skip_verify_sync(self) -> bool: + return self.pre_draft_step == 0 + + +class BaseMtpPlanner(ABC): + """定义 SpecEngine 与不同 MTP 规划器之间的统一调用接口。""" + + @abstractmethod + def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: + """为当前 decode 迭代生成执行计划。 + + Args: + decode_reqs: 当前参与 decode 的逻辑请求列表。规划器可以读取 + 请求的输出进度,判断请求是否已经持有上一轮生成的 + draft proposal;DP 空 rank 对应空列表。 + origin_batch_size: 进入动态压缩前的物理 verify 行数。 + + Returns: + 本轮 target verify 使用的动态 batch size、下一轮需要生成的 draft + step,以及描述当前 proposal 布局的上一轮 draft step。 + """ + + raise NotImplementedError + + @abstractmethod + def update_statics( + self, + plan: SpecDecodePlan, + proposal: SpecProposal, + req_num: int, + accept_lengths, + ) -> None: + """在本轮 verify 完成后更新规划器的运行时统计。 + + Args: + plan: 本轮 decode 实际采用的执行计划,用于确定被验证的配置。 + proposal: proposer 生成的模式专属输出。需要额外调度信息的规划器 + 直接读取自己的 proposal 子类,其他规划器忽略该对象。 + req_num: 本轮逻辑请求数量。 + accept_lengths: 每个请求本轮提交的 token 数量,包含必然提交的 + target token。 + """ + + raise NotImplementedError + + +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 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: + 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] + + +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..e7a6828d66 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/dspark.py @@ -0,0 +1,177 @@ +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 ( + BaseMtpPlanner, + SpecDecodePlan, + _InferCostMsTable, +) +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 + + +class DSparkPlanner(BaseMtpPlanner): + """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, 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) + 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, origin_batch_size: int) -> SpecDecodePlan: + req_num = len(decode_reqs) + full_batch_size = origin_batch_size + dynamic_batch_size = full_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( + 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, + ) + + def update_statics( + self, + plan: SpecDecodePlan, + proposal: DSparkSpecProposal, + req_num: int, + accept_lengths, + ) -> None: + if proposal.schedule_scores_cpu is not None: + self._update_confidence_probs( + confidence_probs=proposal.schedule_scores_cpu, + 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) + 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_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 + + 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 + + # 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), + 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.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, + ) + 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..b46fdd1ceb --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/fixed.py @@ -0,0 +1,34 @@ +from __future__ import annotations + +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.""" + + def __init__(self, max_draft_step: int) -> None: + self.max_draft_step = int(max_draft_step) + + def plan(self, decode_reqs: List, origin_batch_size: int) -> SpecDecodePlan: + return SpecDecodePlan( + 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, + ) + + def update_statics( + self, + plan: SpecDecodePlan, + proposal: SpecProposal, + req_num: int, + accept_lengths, + ) -> 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 new file mode 100644 index 0000000000..b25c62b9b7 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/planner/lightspec.py @@ -0,0 +1,341 @@ +from __future__ import annotations + +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, + SpecDecodePlan, + _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 + from lightllm.server.router.model_infer.mtp_speculative.proposers.base import SpecProposal + + +class LightSpecPlanner(BaseMtpPlanner): + """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 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, + backend: ModeBackend, + ) -> None: + self.spec_mode = spec_mode + 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 + # 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 + + # 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 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. + 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(origin_batch_size, available_batch_size) + all_reqs_have_proposals = req_num_with_proposals == req_num + + 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 + # J(N, B, d) has a progress observation. + 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 + 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) + + 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=plan.origin_batch_size, + dynamic_batch_size=plan.dynamic_batch_size, + draft_step=draft_step, + pre_draft_step=plan.pre_draft_step, + all_reqs_have_proposals=plan.all_reqs_have_proposals, + ) + + def update_statics( + self, + plan: SpecDecodePlan, + proposal: SpecProposal, + req_num: int, + accept_lengths, + ) -> 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 _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)) + # 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,) + 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, + 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) + + # 只有完整 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)) + 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.estimate(dynamic_batch_size) + self._get_draft_cost_ms( + 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/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py new file mode 100644 index 0000000000..dd0d8d05ac --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/__init__.py @@ -0,0 +1,45 @@ +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: "ModeBackend", enable_dynmaic_mtp: bool) -> "BaseSpecProposer": + if spec_mode == "dspark": + 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.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.mtp_speculative.proposers.eagle3 import Eagle3Proposer + + return Eagle3Proposer(backend=backend, enable_dynmaic_mtp=enable_dynmaic_mtp) + if spec_mode == "eagle_with_att": + from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer + + 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 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}") + + +__all__ = [ + "build_spec_proposer", +] diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py new file mode 100644 index 0000000000..4b3d0ed6c6 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, List, Optional + +import torch + +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput + +if TYPE_CHECKING: + 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 `[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. + """ + + token_ids: torch.Tensor + extra_mem_indexes_cpu: List[MtpMemIndexesToFree] = field(default_factory=list) + + +class BaseSpecProposer(ABC): + """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 + the prefill-state and proposal hooks. The proposer returns candidate ids + but does not verify acceptance; verification is handled by SpecEngine. + """ + + 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( + self, + target_model_input: ModelInput, + target_model_output: ModelOutput, + target_next_token_ids: torch.Tensor, + ) -> None: + """Build draft KV/state from target prefill before the first decode verify. + + Inputs: + - `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. + - `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`. + """ + + raise NotImplementedError + + @abstractmethod + 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] + draft_step: int, + accept_len: Optional[torch.Tensor] = None, # [req_num] + ) -> SpecProposal: + """Generate candidate tokens after one target decode forward. + + `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. + + 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 new file mode 100644 index 0000000000..23f81f546a --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dflash.py @@ -0,0 +1,165 @@ +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, + MtpMemIndexesToFree, +) +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 + + +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 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 + 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( + 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, + ) -> DFlashSpecProposal: + """提交 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) + + 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() + + # 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() + ) + 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) + 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: + 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, + 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 new file mode 100644 index 0000000000..5e3ea5694f --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -0,0 +1,190 @@ +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.pin_mem_manager import g_pin_mem_manager +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 ( + DSparkSpecProposal, +) + + +class DSparkProposer(BaseSpecProposer): + """DSpark semi-autoregressive parallel-block proposer. + + 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. + """ + + 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 + 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( + 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, + ) -> DSparkSpecProposal: + """提交 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() + + # 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() + ) + 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 + 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 = ( + confidence_logits[:, :draft_step] + .sigmoid() + .clamp( + min=0.01, + max=0.99, + ) + .contiguous() + ) + + 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 DSparkSpecProposal( + token_ids=proposal_token_ids, + 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/eagle3.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py new file mode 100644 index 0000000000..4ce867ff16 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle3.py @@ -0,0 +1,19 @@ +import torch + +from lightllm.common.basemodel.batch_objs import ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer + + +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: + 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): + 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 new file mode 100644 index 0000000000..629eb7fd3e --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_no_att.py @@ -0,0 +1,108 @@ +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.proposers.base import BaseSpecProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import EagleSpecProposal + + +class EagleNoAttProposer(BaseSpecProposer): + """不使用 attention KV cache 的 EAGLE 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, + ) -> EagleSpecProposal: + 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, + 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 = 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/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..3d2c0a0e86 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -0,0 +1,227 @@ +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.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, +) +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 EagleWithAttProposer(BaseSpecProposer): + """使用 attention KV cache 的 EAGLE 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 + 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, + 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: + """提交验证结果对应的 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 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,) + 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._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._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, + accept_len=accept_len, + 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._gen_argmax_token_ids_and_prob(draft_output) + schedule_scores_by_step.append(draft_token_probs.float().unsqueeze(1)) + else: + 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) + + 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 _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, + 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/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..ee70a26616 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/proposal_type.py @@ -0,0 +1,34 @@ +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 + + +@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/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..ea6a2d5e11 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_no_att.py @@ -0,0 +1,100 @@ +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.proposers.base import BaseSpecProposer +from lightllm.server.router.model_infer.mtp_speculative.proposers.proposal_type import VanillaSpecProposal + + +class VanillaNoAttProposer(BaseSpecProposer): + """不使用 attention KV cache 的 Vanilla chained MTP 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 + # Vanilla No-Att 不维护 KV cache。每一级 draft 只需要每个请求本轮 + # 最后接受位置的 token 和 hidden,因此先把 target verify 布局从 + # [verify_batch_size, ...] 压缩为 [req_num, ...],后续所有 draft + # model 都只对这 req_num 行进行推理。 + 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 = 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)] + + 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/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..ac844cc8d6 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/vanilla_with_att.py @@ -0,0 +1,185 @@ +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.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 + + +class VanillaWithAttProposer(BaseSpecProposer): + """使用 attention KV cache 的 Vanilla chained MTP 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 + + # 在局部副本上逐级左移 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, + 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: + """运行完整 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 = 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, + 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="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/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py new file mode 100644 index 0000000000..935fe92d33 --- /dev/null +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -0,0 +1,139 @@ +from __future__ import annotations + +from collections import Counter +from typing import TYPE_CHECKING, List, Tuple + +import torch + +from lightllm.common.basemodel.triton_kernel.mtp_utils import ( + linear_att_mtp_state_index_update, + mtp_scatter_next_token_ids, + mtp_verify, +) + +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 ( + MtpMemIndexesToFree, + SpecProposal, + ) + + +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, + b_req_idx: torch.Tensor, + b_req_mtp_start_loc: torch.Tensor, + b_mtp_index: torch.Tensor, +) -> 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: + 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 scatter_mtp_next_tokens( + backend: ModeBackend, + 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.""" + + 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( + req_to_next_token_ids=sampling_params_manager.req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + 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=( + 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: ModeBackend, + decode_reqs: List[InferReq], + accept_lengths_cpu: torch.Tensor, + verify_run_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_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_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) + + +def free_mem_indexes( + backend: ModeBackend, + extra_mem_indexes_cpu: List[MtpMemIndexesToFree], +) -> None: + """Free all KV indexes described by the unified MTP memory list.""" + + mem_indexes_to_free = [] + 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)) + + +__all__ = [ + "alloc_mem_indexes", + "free_mem_indexes", + "record_request_mtp_metrics", + "scatter_mtp_next_tokens", + "verify_mtp_tokens", +] 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/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/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index fc6ed9f059..4fed9509a9 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -217,6 +217,15 @@ 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_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)) @@ -242,14 +251,53 @@ 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() + 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}") + + +@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) + # 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. + 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/lightllm/utils/profile_max_tokens.py b/lightllm/utils/profile_max_tokens.py index e3a62b62ea..5ec2a9145e 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,55 @@ 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 + + # 不同 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 + + +@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 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/benchmark/static_inference/static_benchmark.py b/test/benchmark/static_inference/static_benchmark.py index b3faf99130..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 @@ -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), @@ -573,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), @@ -580,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()) @@ -617,66 +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()), - 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(), 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/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 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..2a5fc17298 --- /dev/null +++ b/test/start_scripts/qwen35/qwen35_pd_1p1d.sh @@ -0,0 +1,123 @@ +#!/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.75 + --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 5 + "${CHAT_TEMPLATE_ARGS[@]}" + --pd_trans_mode nccl + --pd_kv_page_size 4096 + --pd_master_ip 127.0.0.1 + --pd_master_port "${PORT}" +) + +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 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/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 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/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py new file mode 100644 index 0000000000..17ddcd98da --- /dev/null +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -0,0 +1,78 @@ +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( + 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_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, + ) + 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( + batch_step_size_before_split=8, + split_batch_size=32, + batch_step_size_after_split=16, + max_batch_size=32, + ) == [ + 8, + 16, + 24, + 32, + ] + + +def test_instance_and_public_static_schedule_match(_graph_args): + 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, + ) + + +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_batch_step_size_after_split_controls_capture_range(_graph_args): + assert _batch_sizes(max_batch_size=8, batch_stride=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..22748b7294 --- /dev/null +++ b/unit_tests/common/basemodel/test_hidden_collector.py @@ -0,0 +1,193 @@ +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, + LayerHiddenCollector, + MtpHeadOutputCollector, + NoopHiddenCollector, +) + + +class _IdentityPreInfer: + @staticmethod + def _tpsp_allgather(input, infer_state): + del infer_state + return input + + +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}, {}), + ) + + +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_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) + + collector0.add_final_hidden(hidden0) + 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_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() + + +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()) + 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) + + collector0.add(layer_index=0, hidden=hidden0) + collector1.add(layer_index=0, hidden=hidden1) + + 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) + + +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.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_output(infer_state=None).spec_hidden + + assert collected.data_ptr() == final_hidden.data_ptr() + + +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) + + 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) + + 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_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 + + +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_output(infer_state=infer_state).spec_hidden + + 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_model_input.py b/unit_tests/common/basemodel/test_model_input.py new file mode 100644 index 0000000000..056513bdde --- /dev/null +++ b/unit_tests/common/basemodel/test_model_input.py @@ -0,0 +1,219 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.batch_objs import ModelInput + + +def _create_model_input(*, is_prefill=False): + batch_size = 2 + kwargs = dict( + 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), + 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 is_prefill: + kwargs["max_cache_len"] = 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) + + +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), + 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": []}], + ) + + +def test_decode_carries_raw_shared_radix_metadata(): + model_input = _create_model_input() + + 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_prefill_does_not_require_shared_radix_metadata(): + model_input = _create_model_input(is_prefill=True) + + 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_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 + + 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, + 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] + + +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, + 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=[], + mtp_draft_input_hiddens=torch.empty((0, 4), dtype=torch.float32), + ) + 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] + assert torch.equal(padded_input.mtp_draft_input_hiddens, torch.zeros((1, 4))) + + +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 new file mode 100644 index 0000000000..6f9477e294 --- /dev/null +++ b/unit_tests/common/basemodel/test_model_output.py @@ -0,0 +1,164 @@ +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 ModelInput, ModelMtpOutputCollector, ModelOutput + + +def test_decode_unpad_slices_spec_output_with_logits(): + model = TpPartBaseModel.__new__(TpPartBaseModel) + output = ModelOutput( + logits=torch.arange(24).view(6, 4), + 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.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.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), + mtp_collector=ModelMtpOutputCollector(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.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/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py new file mode 100644 index 0000000000..c7ad04f56b --- /dev/null +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -0,0 +1,128 @@ +from types import SimpleNamespace + +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, + MtpHeadOutputCollector, + NoopHiddenCollector, +) +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, + 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) + + +@pytest.mark.parametrize( + "spec_mode,is_draft_model,expected", + [ + (None, False, 1), + ("eagle3", False, 8), + ("eagle3", True, 1), + ("vanilla_with_att", True, 1), + ("vanilla_no_att", True, 1), + ("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 + + +@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 + + +@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) + + 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, MtpHeadOutputCollector), + ], +) +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_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/test_prefill_cuda_graph_state.py b/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py new file mode 100644 index 0000000000..bed4f3967f --- /dev/null +++ b/unit_tests/common/basemodel/test_prefill_cuda_graph_state.py @@ -0,0 +1,81 @@ +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.input_ids = torch.empty(4, dtype=torch.int64) + 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( + input_ids=torch.empty(4, dtype=torch.int64), + 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_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/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..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/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}" ) 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 new file mode 100644 index 0000000000..2cb95f5923 --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_dynamic_mtp_utils.py @@ -0,0 +1,371 @@ +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() + + 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_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, + 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_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) + ) + # 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_preserves_shared_radix_metadata(): + 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_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"), + 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, + ) + 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_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: + 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 = scores[req_idx, : max_draft_step + 1].clone() + row[0] = 1.0 + row[1:] = torch.clamp(row[1:], min=0.01, max=0.99) + scores[req_idx, : max_draft_step + 1] = torch.cumprod(row, dim=0) + return scores + + +def _flat_cumprod_scores( + b_req_idx: torch.Tensor, + req_to_next_token_scores: torch.Tensor, + max_draft_step: int, +) -> torch.Tensor: + 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_scores = [] + for offset in range(all_num): + req_idx = int(b_req_idx[offset].item()) + mtp_index = offset % (max_draft_step + 1) + flat_scores.append(scores[req_idx, mtp_index]) + return torch.stack(flat_scores) + + +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_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_scores(req_num: int, max_draft_step: int, rows): + max_req = req_num + scores = torch.zeros((max_req + 1, 16), dtype=torch.float32, device="cuda") + for req_idx, row in enumerate(rows): + 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 scores, b_req_idx + + +@pytest.mark.parametrize("max_draft_step", [1, 3]) +def test_cumprod_scores_kernel(max_draft_step: int): + req_num = 2 + scores, b_req_idx = _make_batch_scores( + req_num, + max_draft_step, + rows=[ + [1.0] + [0.5] * max_draft_step, + [1.0] + [0.2] * max_draft_step, + ], + ) + 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_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_scores_clamps_invalid_values(): + max_draft_step = 2 + req_num = 1 + 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 = 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_scores_clamps_boundary_values(): + max_draft_step = 3 + req_num = 1 + 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_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 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) + assert row[3].item() == pytest.approx(0.99 * 0.01 * 0.5, rel=1e-4) + + +def test_sample_select_count(): + max_draft_step = 3 + req_num = 3 + scores, b_req_idx = _make_batch_scores( + req_num, + 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 * (max_draft_step + 1) + for dynamic_batch_size in [3, 8, all_num]: + 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(), + max_draft_step=max_draft_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(): + max_draft_step = 3 + scores, b_req_idx = _make_batch_scores( + 3, + 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_row_mask( + dynamic_batch_size=np.int64(8), + b_req_idx=b_req_idx, + req_to_next_token_scores=scores, + max_draft_step=np.int64(max_draft_step), + ) + assert int(select.sum().item()) == 8 + + +def test_sample_topk_by_cumprod_score(): + max_draft_step = 3 + scores, b_req_idx = _make_batch_scores( + 3, + 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_scores = _flat_cumprod_scores(b_req_idx, scores, max_draft_step) + for dynamic_batch_size in [1, 4, 8, 12]: + 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(), + max_draft_step=max_draft_step, + ) + _assert_topk_mask(select, flat_scores, dynamic_batch_size) + + +def test_sample_picks_highest_cumprod_rows(): + max_draft_step = 1 + scores, b_req_idx = _make_batch_scores( + 2, + max_draft_step, + rows=[ + [1.0, 0.9], + [1.0, 0.1], + ], + ) + flat_scores = _flat_cumprod_scores(b_req_idx, scores, max_draft_step) + select = sample_dynamic_mtp_row_mask( + dynamic_batch_size=2, + b_req_idx=b_req_idx, + req_to_next_token_scores=scores.clone(), + max_draft_step=max_draft_step, + ) + _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 + + +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_mtp_row_mask( + dynamic_batch_size=2, + b_req_idx=b_req_idx, + req_to_next_token_scores=scores.clone(), + max_draft_step=max_draft_step, + ) + _assert_topk_mask(select, flat_scores, 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..c1ef686f8e --- /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_spec_fa3_decode_params + + +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_mtp_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_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_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_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_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_mtp_shared_group=b_mark_mtp_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_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_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_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_mtp_shared_group=b_mark_mtp_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_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/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..f6f9df9d44 --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_mtp_utils.py @@ -0,0 +1,101 @@ +import pytest +import torch + +if not torch.cuda.is_available(): + pytest.skip("requires CUDA", allow_module_level=True) + +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]], + 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") + 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.7, 0.6]], + dtype=torch.float32, + 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=req_to_next_token_ids, + b_req_mtp_start_loc=b_req_mtp_start_loc, + 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, + 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(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], [2, 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"), + 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, + 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/qwen2_vl/test_infer_struct.py b/unit_tests/models/qwen2_vl/test_infer_struct.py new file mode 100644 index 0000000000..b239e64067 --- /dev/null +++ b/unit_tests/models/qwen2_vl/test_infer_struct.py @@ -0,0 +1,68 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.basemodel.infer_struct import InferStateInfo +from lightllm.models.qwen2_vl.infer_struct import Qwen2VLInferStateInfo + + +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_positions) + + +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), + ) + + +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_prefill_rejects_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 + + 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/models/test_qwen3_dspark_model_output.py b/unit_tests/models/test_qwen3_dspark_model_output.py new file mode 100644 index 0000000000..1dfdfeae23 --- /dev/null +++ b/unit_tests/models/test_qwen3_dspark_model_output.py @@ -0,0 +1,436 @@ +from types import SimpleNamespace + +import pytest +import torch + +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 +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 +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) + output = ModelOutput( + logits=torch.arange(48).view(12, 4), + 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), + ), + ) + + unpadded = model._create_unpad_decode_model_output(output, origin_batch_size=8) + + assert isinstance(unpadded, ModelOutput) + assert unpadded.logits.shape == (8, 4) + 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.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)), + mtp_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.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.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)) + + +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", + [ + {"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 + self.quant_cfg = quant_cfg + + class RecordingROWMMWeight: + def __init__(self, **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=network_config, + quant_cfg=quant_cfg, + ) + + 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_adapter_uses_current_checkpoint_rope_layout(monkeypatch): + def init_dspark_config(self): + self.config = { + "dflash_config": {"mask_token_id": 1}, + "rope_parameters": { + "rope_theta": 1_000_000, + "factor": 32.0, + "original_max_position_embeddings": 8192, + "rope_type": "yarn", + }, + } + + monkeypatch.setattr(Qwen3DSparkModel, "_init_config", init_dspark_config) + model = Qwen3_5DSparkModel.__new__(Qwen3_5DSparkModel) + + model._init_config() + + 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 + 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) + + +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/models/test_qwen3_eagle_model_input.py b/unit_tests/models/test_qwen3_eagle_model_input.py new file mode 100644 index 0000000000..1395153351 --- /dev/null +++ b/unit_tests/models/test_qwen3_eagle_model_input.py @@ -0,0 +1,100 @@ +from types import SimpleNamespace + +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( + 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/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..fcdf4e3295 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -0,0 +1,487 @@ +from contextlib import nullcontext +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_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 ( + FixedSpecPlanner, + LightSpecPlanner, + SpecDecodePlan, +) +from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( + SpecProposal, +) + + +def test_dp_backend_reuses_common_engine_outside_overlap(): + args = SimpleNamespace(mtp_mode="eagle3", mtp_dynamic_verify=False, dp=1) + 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.dp_size = 1 + 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 SpecEngine + 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) + + +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_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.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) + 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 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(): + 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, 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() + + 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 + + +@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=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 = [] + + 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)) + + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + 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"), + ) + 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"), + ) + + 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_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_ids0"].new_empty((3, 7))) + + engine = DPOverlapSpecEngine.__new__(DPOverlapSpecEngine) + engine.proposer = _Proposer() + model_input0 = SimpleNamespace(batch_size=8) + model_input1 = SimpleNamespace(batch_size=16) + model_output0 = SimpleNamespace() + 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_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, + accept_len1=accept_len1, + 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_ids0"] is target_next_token_ids0 + assert calls["target_next_token_ids1"] is target_next_token_ids1 + 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) + + +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" + 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 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_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["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)) + + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.decode_draft_engine = _OverlapEngine() + backend.model = SimpleNamespace( + microbatch_overlap_decode=lambda input0, input1: (model_output0, model_output1), + ) + 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"), + ) + 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=[]) + + assert calls == [ + "plan", + "prepare", + "propose", + "post_wait", + "forward_wait", + "free", + "pre_post", + ] 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..62296634a9 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_generic_pre_process.py @@ -0,0 +1,172 @@ +from types import SimpleNamespace + +import torch +from lightllm.server.router.model_infer.mode_backend import 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 _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) + + 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_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), + ] + + ( + 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, + 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 + 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, + 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 == [] + 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_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_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_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py new file mode 100644 index 0000000000..2e24262e7e --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -0,0 +1,400 @@ +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 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, +) +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, +) + + +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_cuda(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), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2))), + ) + for model_input in (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( + ModelOutput( + 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) + ) + + def map_draft_vocab_to_main_vocab(self, token_ids): + return token_ids + + +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, 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, + is_prefill=False, + multimodal_params=[{"images": [], "audios": []}] * batch_size, + ) + + +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, + 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), + ) + 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((3, 1)), + 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), + accept_len1=torch.tensor([1, 3], dtype=torch.int32), + draft_step=2, + ) + + assert draft_model.extend_batch_sizes is None + 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 proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None + 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_supports_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=0), + target_model_output0=ModelOutput( + logits=torch.empty((0, 1)), + 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), + accept_len1=torch.empty((0,), dtype=torch.int32), + draft_step=2, + ) + + assert proposal.token_ids.shape == (0, 2) + assert draft_model.decode_batch_sizes == [(0, 0), (0, 0)] + + +def test_overlap_eagle_returns_dynamic_schedule_scores(monkeypatch): + _patch_cpu_req_start_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_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), + 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), + ), + 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), + accept_len1=torch.tensor([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), + 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), + accept_len1=torch.tensor([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): + _patch_cpu_req_start_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 = DpOverlapEagle3Proposer(backend=backend, enable_dynmaic_mtp=False) + monkeypatch.setattr( + mtp_utils, + "alloc_mem_indexes", + lambda token_count: torch.arange(token_count, dtype=torch.int32), + ) + 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((3, 1)), + 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), + accept_len1=torch.tensor([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 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 == [(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 proposal.extra_mem_indexes_cpu[0].free_mask_cpu is None + + +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])) + + +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_eagle_with_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py new file mode 100644 index 0000000000..f50780d53c --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py @@ -0,0 +1,224 @@ +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.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(): + 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 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 new file mode 100644 index 0000000000..7baf34061b --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_planner.py @@ -0,0 +1,992 @@ +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine +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, + FixedSpecPlanner, + 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.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.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 import utils as mtp_utils + + +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=enable_decode_microbatch_overlap, + 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, + backend=backend, + ) + + +def build_dspark_planner(max_draft_step: int = 3, block_size: int = 3): + backend = SimpleNamespace( + max_draft_step=max_draft_step, + 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): + 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.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)], + ) + return engine._build_mtp_planner( + spec_mode=spec_mode, + enable_dynmaic_mtp=enable_dynmaic_mtp, + ) + + +def test_fixed_planner_returns_static_plan(): + planner = FixedSpecPlanner(max_draft_step=3) + plan = planner.plan(decode_reqs=build_decode_reqs(4), origin_batch_size=16) + + assert isinstance(planner, BaseMtpPlanner) + 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_common_engine_accepts_empty_fixed_decode_batch(): + engine = SpecEngine.__new__(SpecEngine) + engine.planner = FixedSpecPlanner(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_common_engine_delegates_empty_dp_batch_to_lightspec_planner(): + engine = SpecEngine.__new__(SpecEngine) + 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=1, + pre_draft_step=2, + ) + + +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 == { + "fill_draft_model_kv_state", + "plan_decode", + "prepare_decode_model_input", + "propose_next", + "update_planner_statics", + } + + +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(2, dtype=torch.int64).view(2, 1), + extra_mem_indexes_cpu=[], + 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), + ) + + 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["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_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 = VanillaSpecProposal( + 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_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(): + 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 == (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) + assert dflash_planner.draft_steps == (3,) + + eagle_planner = build_planner("eagle3") + 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=[ + 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_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, + VanillaNoAttProposer, + EagleWithAttProposer, + EagleNoAttProposer, + DFlashProposer, + DSparkProposer, + ) + dp_overlap_proposer_types = ( + DpOverlapVanillaWithAttProposer, + DpOverlapVanillaNoAttProposer, + DpOverlapEagleWithAttProposer, + DpOverlapEagleNoAttProposer, + ) + + for proposer_type in proposer_types: + assert proposer_type.__bases__ == (BaseSpecProposer,) + 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(): + 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_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 = EagleNoAttProposer( + backend=SimpleNamespace(draft_models=[]), + enable_dynmaic_mtp=True, + ) + next_token_ids = torch.tensor([10, 11], dtype=torch.int64) + + proposal = proposer.propose_next( + 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, + ) + + assert isinstance(proposal, EagleSpecProposal) + assert proposal.token_ids.shape == (2, 0) + assert proposal.schedule_scores.shape == (2, 0) + + +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()))) + ) + ) + mtp_utils.free_mem_indexes( + backend=backend, + 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 + assert freed[0].tolist() == [11, 12, 13, 20, 22] + + +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 + self.mtp_step = mtp_step + self.accepted = 0 + self.verified = 0 + self.verify_steps = 0 + + def update_mtp_accepted_token_num(self, accept_token_num: int): + self.accepted += accept_token_num + + def update_mtp_verify_token_num(self, verify_token_num: int): + self.verified += verify_token_num + + def update_mtp_verify_step_num(self, verify_step_num: int): + self.verify_steps += verify_step_num + + backend = SimpleNamespace(is_master_in_dp=True) + req0 = MetricReq(req_idx=10, mtp_step=3) + req1 = MetricReq(req_idx=11, mtp_step=3) + + mtp_utils.record_request_mtp_metrics( + backend=backend, + decode_reqs=[req0, req1], + accept_lengths_cpu=torch.tensor([2, 1], dtype=torch.int32), + verify_run_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) + mtp_utils.record_request_mtp_metrics( + backend=backend, + decode_reqs=[fixed_req], + accept_lengths_cpu=torch.tensor([3], dtype=torch.int32), + verify_run_reqs=[fixed_req] * 4, + ) + 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( + decode_reqs=build_decode_reqs(2), + 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 + + +def test_lightspec_collects_full_width_progress_before_adapting(): + planner = build_lightspec_planner(spec_mode="eagle_with_att") + 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)) + + 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 + + +def test_lightspec_selects_eagle_draft_depth_and_verify_capacity(): + 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.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) + for _ in range(8): + 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_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)): + 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, + 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.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], + req_num=2, + dynamic_batch_size=16, + verified_draft_step=7, + ) + + 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 + + +def test_lightspec_bounds_verify_to_existing_proposals(): + planner = build_lightspec_planner() + + 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 + assert ready_batch.dynamic_batch_size == 8 + 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_lets_planner_count_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.all_reqs_have_proposals + assert ready_plan.dynamic_batch_size == 8 + 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 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, + ) + + 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( + dynamic_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): + 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, + 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 == (1, 2, 3) + assert plan.draft_step >= 1 + + +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( + req_num=4, + verify_batch_size=8, + draft_step=3, + ) + + 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_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( + req_num=2, + verify_batch_size=8, + draft_step=3, + ) + + assert draft_cost_ms == 1.5 + + +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, + spec_mode=spec_mode, + ) + 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( + req_num=8, + verify_batch_size=16, + draft_step=7, + ) + + assert draft_cost_ms == 64.0 + + +def test_block_planner_prices_commit_and_complete_block(): + planner = build_lightspec_planner( + max_draft_step=7, + spec_mode="dflash", + block_size=7, + ) + planner.draft_infer_costs.update(batch_size=8, infer_cost_ms=0.4) + + draft_cost_ms = planner._get_draft_cost_ms( + req_num=2, + verify_batch_size=16, + draft_step=7, + ) + + 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() + + 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 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 + + +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_estimates_unseen_shapes_from_prefix_survival(): + 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_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(spec_mode="eagle_with_att") + 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( + accept_lengths=[4, 4], + req_num=2, + dynamic_batch_size=8, + verified_draft_step=3, + ) + planner.pre_draft_step = 1 + + 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 + assert plan.draft_step == 3 + + +def test_engine_skips_feedback_for_a_mixed_proposal_batch(): + engine = SpecEngine.__new__(SpecEngine) + engine.planner = build_lightspec_planner() + plan = SpecDecodePlan( + origin_batch_size=8, + dynamic_batch_size=5, + draft_step=3, + pre_draft_step=3, + all_reqs_have_proposals=False, + ) + + engine.update_planner_statics( + plan=plan, + proposal=SpecProposal( + token_ids=torch.empty((0,), dtype=torch.int64), + extra_mem_indexes_cpu=[], + ), + 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 = build_dspark_planner() + 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) + 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 = DSparkSpecProposal( + token_ids=torch.empty((0,), dtype=torch.int64), + extra_mem_indexes_cpu=[], + schedule_scores_cpu=torch.from_numpy(confidence_probs), + ) + engine = SpecEngine.__new__(SpecEngine) + engine.planner = planner + + 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_statics( + 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), origin_batch_size=8) + + assert first_plan.dynamic_batch_size == 8 + assert second_plan.dynamic_batch_size == 4 + 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.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) + + 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 = DSparkPlanner._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/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..8ce9114aab --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_no_att.py @@ -0,0 +1,220 @@ +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.dp_overlap_proposers.vanilla_no_att import ( + DpOverlapVanillaNoAttProposer, +) +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_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" + 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.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] + + 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], 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]])) + 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.cpu(), 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) + + +def test_vanilla_no_att_fill_hooks_are_noops(): + backend = SimpleNamespace(draft_models=[]) + proposer = VanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False) + overlap_proposer = DpOverlapVanillaNoAttProposer(backend=backend, enable_dynmaic_mtp=False) + + proposer.fill_draft_model_kv_state(None, None, None) + overlap_proposer.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_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py new file mode 100644 index 0000000000..b98152c855 --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_overlap.py @@ -0,0 +1,148 @@ +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.dp_overlap_proposers.vanilla_no_att import ( + DpOverlapVanillaNoAttProposer, +) +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_cuda(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, + 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) + ) + + +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, + accept_len0=torch.ones((1,), dtype=torch.int32), + target_model_input1=None, + target_model_output1=None, + target_next_token_ids1=target_next_token_ids1, + accept_len1=torch.ones((1,), 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" + 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) + 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, + target_model_output0=ModelOutput( + 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), + 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), + accept_len1=torch.tensor([1], dtype=torch.int32, device=device), + draft_step=2, + ) + + assert proposal.token_ids.tolist() == [ + [1, 1], + [0, 0], + ] + assert proposal.extra_mem_indexes_cpu == [] + 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") +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) + 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((0, 1), device=device), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((0, 2), device=device)), + ) + model_output1 = ModelOutput( + 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), + 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), + accept_len1=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 == [(0, 0)] + assert draft_models[1].decode_batch_sizes == [(0, 0)] 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..c912d81212 --- /dev/null +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_vanilla_prefill.py @@ -0,0 +1,238 @@ +from types import SimpleNamespace + +import pytest +import torch + +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_cuda(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) + + 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), + 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)) + + +@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), + ) 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) 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() diff --git a/unit_tests/utils/test_speculative_utils.py b/unit_tests/utils/test_speculative_utils.py new file mode 100644 index 0000000000..aaa3e05268 --- /dev/null +++ b/unit_tests/utils/test_speculative_utils.py @@ -0,0 +1,386 @@ +import json +from importlib import import_module +from types import SimpleNamespace + +import pytest +import torch + +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.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.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), + ("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), + ], +) +def test_attention_backend_selects_dynamic_spec_layout( + monkeypatch, + mtp_mode, + is_draft_model, + draft_step, + dynamic_spec, + expected, +): + 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, + mtp_manager=SimpleNamespace(get_decode_draft_step=lambda _: draft_step), + ) + ) + + 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.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", + [ + ("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(model_type, spec_mode, expected_class_name): + model_class = get_draft_model_class( + model_cfg={"model_type": model_type}, + spec_mode=spec_mode, + ) + + assert model_class.__name__ == expected_class_name + + +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( + "model_type, spec_mode", + [ + ("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_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": model_type}, + spec_mode=spec_mode, + ) + + +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", + get_config_dict, + ) + 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) + + collector = hidden_collector_module.LayerHiddenCollector(model=model) + new_collector = collector.new_instance() + + assert collector.layer_ids == frozenset((1, 20, 36)) + assert new_collector.layer_ids == collector.layer_ids + 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})) + + 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_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})) + + 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