From 68d2418996c421fbaf351796671658e095730fad Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 21 Sep 2026 13:59:07 +0800 Subject: [PATCH 01/11] fix(mtp): preserve speculative groups in TP-SP batch alignment --- lightllm/common/basemodel/basemodel.py | 17 ++++-- lightllm/common/basemodel/cuda_graph.py | 5 +- .../basemodel/test_cuda_graph_layout.py | 59 +++++++++++++++++++ .../common/basemodel/test_model_output.py | 59 +++++++++++++++++++ 4 files changed, 134 insertions(+), 6 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 7d1551e053..f2e3ce6be0 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -3,6 +3,7 @@ # os.environ["CUDA_LAUNCH_BLOCKING"] = "1" import gc import copy +import math import json import torch import torch.nn.functional as F @@ -102,6 +103,14 @@ def __init__(self, kvargs): self.mem_fraction = kvargs.get("mem_fraction", 0.9) self.tp_world_size_ = get_dp_world_size() self.enable_tpsp_mix_mode = get_env_start_args().enable_tpsp_mix_mode + # Fixed speculative batches must contain whole verify groups even after TP/SP padding. + self.decode_batch_alignment = math.lcm( + self.mtp_manager.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model), + self.tp_world_size_ if self.enable_tpsp_mix_mode else 1, + ) + self.graph_max_batch_size = ( + triton.cdiv(self.graph_max_batch_size, self.decode_batch_alignment) * self.decode_batch_alignment + ) self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) self.prefill_graph: PrefillCudaGraph = None @@ -592,11 +601,9 @@ def _decode( ) origin_batch_size = model_input.batch_size - # 空 DP rank 先补出一个 dummy request;TPSP 模式下继续将 batch size - # 向上对齐到 TP world size 的整数倍,保证后续切分得到合法 shape。 + # 空 DP rank 也需要完整的 dummy verify group;TP/SP 切分不能破坏 MTP 分组。 infer_batch_size = max(1, origin_batch_size) - if self.args.enable_tpsp_mix_mode: - infer_batch_size = triton.cdiv(infer_batch_size, self.tp_world_size_) * self.tp_world_size_ + infer_batch_size = triton.cdiv(infer_batch_size, self.decode_batch_alignment) * self.decode_batch_alignment # CUDA Graph 可能继续向上对齐 batch size,并因此加入 seq_len=2 的 # dummy request。先用最终可能出现的 KV 长度判断 graph,再统一 padding 一次。 @@ -819,7 +826,7 @@ def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1 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_ + infer_batch_size = triton.cdiv(infer_batch_size, self.decode_batch_alignment) * self.decode_batch_alignment 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) diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index dd15af99f5..ee888c1a00 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -3,6 +3,7 @@ import torch.distributed as dist import copy import bisect +import math import triton from typing import Optional from lightllm.utils.log_utils import init_logger @@ -47,7 +48,9 @@ def gen_cuda_graph_batch_sizes( batch_sizes = sorted({size for size in batch_sizes if size < max_batch_size} | {max_batch_size}) if args.enable_tpsp_mix_mode: - batch_sizes = sorted({triton.cdiv(size, tp_world_size) * tp_world_size for size in batch_sizes}) + # Keep complete fixed-layout MTP groups as well as TP/SP shards. + alignment = math.lcm(batch_step_size_before_split, tp_world_size) + batch_sizes = sorted({triton.cdiv(size, alignment) * alignment for size in batch_sizes}) assert batch_sizes[-1] == max_batch_size return batch_sizes diff --git a/unit_tests/common/basemodel/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py index 17ddcd98da..bf6e52db6e 100644 --- a/unit_tests/common/basemodel/test_cuda_graph_layout.py +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -3,7 +3,11 @@ import pytest import lightllm.common.basemodel.cuda_graph as cuda_graph_module +import lightllm.common.basemodel.basemodel as basemodel_module +import lightllm.common.basemodel.mtp_manager as mtp_manager_module +from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.common.basemodel.cuda_graph import CudaGraph +from lightllm.common.basemodel.mtp_manager import MtpManager @pytest.fixture(autouse=True) @@ -76,3 +80,58 @@ def test_batch_step_size_after_split_controls_capture_range(_graph_args): 42, 56, ] + + +@pytest.mark.parametrize("tp_size", [1, 2, 4, 8]) +@pytest.mark.parametrize("mtp_step", [1, 2]) +@pytest.mark.parametrize("dynamic,is_draft", [(False, False), (True, False), (False, True)]) +@pytest.mark.parametrize("overlap", [False, True]) +def test_mtp_tpsp_layout(monkeypatch, _graph_args, tp_size, mtp_step, dynamic, is_draft, overlap): + args = _graph_args + args.enable_tpsp_mix_mode = True + args.enable_decode_microbatch_overlap = overlap + args.mtp_mode = "eagle_with_att" + args.mtp_step = mtp_step + args.mtp_dynamic_verify = dynamic + args.mem_fraction = 0.8 + monkeypatch.setattr(basemodel_module, "get_env_start_args", lambda: args) + monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) + monkeypatch.setattr(MtpManager, "_instance", None) + monkeypatch.setattr(basemodel_module, "get_dp_world_size", lambda: tp_size) + monkeypatch.setattr(basemodel_module, "get_llm_data_type", lambda: None) + monkeypatch.setattr(basemodel_module, "TorchMemorySaverWrapper", lambda _: None) + + class StopBeforeWeights(Exception): + pass + + def stop_init(self): + raise StopBeforeWeights + + monkeypatch.setattr(TpPartBaseModel, "_init_config", stop_init) + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.is_mtp_draft_model = is_draft + with pytest.raises(StopBeforeWeights): + model.__init__(dict(run_mode="normal", weight_dir="unused", max_total_token_num=128, graph_max_batch_size=7)) + + width = 1 if dynamic or is_draft else mtp_step + 1 + assert model.decode_batch_alignment % width == 0 + assert model.decode_batch_alignment % tp_size == 0 + if width == 1: + assert model.decode_batch_alignment == tp_size + assert model.max_req_num == 1000 # Physical padding does not change request capacity. + assert model.graph_max_batch_size % model.decode_batch_alignment == 0 + logical_max = 7 // 2 if overlap else 7 + physical_max = logical_max * (1 if is_draft else mtp_step + 1) + assert physical_max <= model.graph_max_batch_size < physical_max + model.decode_batch_alignment + + sizes = CudaGraph.gen_cuda_graph_batch_sizes( + batch_step_size_before_split=width, + split_batch_size=4 * width, + batch_step_size_after_split=2 * width, + max_batch_size=model.graph_max_batch_size, + tp_world_size=tp_size, + ) + assert sizes[-1] == model.graph_max_batch_size + assert all(size % width == 0 and size % tp_size == 0 for size in sizes) + if tp_size == 8 and width == 3: + assert sizes == [24] diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index 4b723e6a1a..b78e61812f 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -1,6 +1,7 @@ from types import SimpleNamespace import torch +import pytest from lightllm.common.basemodel import basemodel from lightllm.common.basemodel.basemodel import TpPartBaseModel @@ -115,6 +116,7 @@ def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): model = TpPartBaseModel.__new__(TpPartBaseModel) model.args = SimpleNamespace(enable_tpsp_mix_mode=enable_tpsp_mix_mode, page_size=1) model.tp_world_size_ = tp_world_size + model.decode_batch_alignment = tp_world_size if enable_tpsp_mix_mode else 1 model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEXES=(99,), page_size=1) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) @@ -159,3 +161,60 @@ def create_infer_state(model_input): assert pad_batch_sizes == [expected_batch_size] assert graph_flags_at_att_init == [need_capture] assert output.logits.shape == (0, 4) + + +@pytest.mark.parametrize("overlap", [False, True]) +@pytest.mark.parametrize("graph_mode", ["eager", "capture", "replay"]) +@pytest.mark.parametrize("rows", [0, 3, 9, 24, 27]) +def test_fixed_mtp_decode_preserves_groups_and_real_rows(overlap, graph_mode, rows): + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.args = SimpleNamespace(enable_tpsp_mix_mode=True) + model.tp_world_size_ = 8 + model.decode_batch_alignment = 24 # TP=8, mtp_step=2. + model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88) + seen_sizes = [] + + def create_state(model_input, microbatch_index=0): + def check_layout(): + assert model_input.batch_size % 3 == 0 + assert model_input.batch_size % 8 == 0 + seen_sizes.append(model_input.batch_size) + + return SimpleNamespace( + b_req_idx=model_input.b_req_idx, + b_seq_len=model_input.b_seq_len, + init_some_extra_state=lambda _: None, + init_att_state=check_layout, + ) + + def forward(state): + return ModelOutput(logits=torch.stack((state.b_req_idx, state.b_seq_len), dim=1)) + + model._create_inferstate = create_state + model._token_forward = forward + model._overlap_tpsp_token_forward = lambda state, infer_state1: (forward(state), forward(infer_state1)) + model.graph = None + if graph_mode != "eager": + model.graph = SimpleNamespace( + can_run=lambda batch_size, max_len_in_batch: batch_size <= 24, + find_closest_graph_batch_size=lambda batch_size: 24, + need_capture=lambda batch_size: graph_mode == "capture", + capture_decode=lambda func, state, **kwargs: func(state, **kwargs), + replay=lambda state, infer_state1=None: ( + forward(state) if infer_state1 is None else (forward(state), forward(infer_state1)) + ), + ) + + model_input = model._create_padded_decode_model_input(_create_empty_decode_input(), rows) + model_input.b_req_idx = torch.arange(rows, dtype=torch.int32) // 3 + model_input.b_seq_len = 10 + torch.arange(rows, dtype=torch.int32) % 3 + expected = torch.stack((model_input.b_req_idx, model_input.b_seq_len), dim=1) + if overlap: + output, empty_output = model._microbatch_overlap_decode_cuda(model_input, _create_empty_decode_input()) + assert empty_output.logits.shape == (0, 2) + else: + output = model._decode(model_input) + + assert seen_sizes == ([48 if rows > 24 else 24] * (2 if overlap else 1)) + torch.testing.assert_close(output.logits, expected) + torch.testing.assert_close(torch.stack((model_input.b_req_idx, model_input.b_seq_len), dim=1), expected) From 7e5926b5c761ad4ddda841057a9aa509b4b4afbd Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 21 Sep 2026 14:05:48 +0800 Subject: [PATCH 02/11] refactor(mtp): centralize decode batch layout initialization --- lightllm/common/basemodel/basemodel.py | 40 +++++++++---------- .../basemodel/test_cuda_graph_layout.py | 23 +++-------- 2 files changed, 24 insertions(+), 39 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index f2e3ce6be0..829950d618 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -84,16 +84,7 @@ def __init__(self, kvargs): self.max_seq_length = kvargs.get("max_seq_length", 1024 * 5) self.return_all_prompt_logics = kvargs.get("return_all_prompt_logics", False) self.data_type = get_llm_data_type() - self.graph_max_batch_size = kvargs.get("graph_max_batch_size", 16) - self.graph_max_batch_size = ( - self.graph_max_batch_size // 2 - if get_env_start_args().enable_decode_microbatch_overlap - else self.graph_max_batch_size - ) self.mtp_manager = MtpManager.get_instance() - self.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) @@ -103,14 +94,7 @@ def __init__(self, kvargs): self.mem_fraction = kvargs.get("mem_fraction", 0.9) self.tp_world_size_ = get_dp_world_size() self.enable_tpsp_mix_mode = get_env_start_args().enable_tpsp_mix_mode - # Fixed speculative batches must contain whole verify groups even after TP/SP padding. - self.decode_batch_alignment = math.lcm( - self.mtp_manager.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model), - self.tp_world_size_ if self.enable_tpsp_mix_mode else 1, - ) - self.graph_max_batch_size = ( - triton.cdiv(self.graph_max_batch_size, self.decode_batch_alignment) * self.decode_batch_alignment - ) + self._init_decode_batch_layout(kvargs.get("graph_max_batch_size", 16)) self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) self.prefill_graph: PrefillCudaGraph = None @@ -282,6 +266,22 @@ def _init_att_backend1(self): self.decode_att_backend1: BaseAttBackend = None return + def _init_decode_batch_layout(self, max_requests: int): + if self.args.enable_decode_microbatch_overlap: + max_requests //= 2 + + mtp = self.mtp_manager + self.decode_batch_alignment = math.lcm( + mtp.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model), + self.tp_world_size_ if self.enable_tpsp_mix_mode else 1, + ) + max_rows = max_requests * mtp.get_decode_batch_multiplier(self.is_mtp_draft_model) + self.graph_max_batch_size = self._align_decode_batch_size(max_rows) + + def _align_decode_batch_size(self, batch_size: int) -> int: + # Preserve complete MTP verify groups and TP/SP shards, including dummy rows. + return triton.cdiv(batch_size, self.decode_batch_alignment) * self.decode_batch_alignment + def _init_cudagraph(self): # When graph covers the configured request length, it must also cover MTP's internal token margin. if self.args.mtp_mode is not None and self.graph_max_len_in_batch >= self.args.max_req_total_len: @@ -602,8 +602,7 @@ def _decode( origin_batch_size = model_input.batch_size # 空 DP rank 也需要完整的 dummy verify group;TP/SP 切分不能破坏 MTP 分组。 - infer_batch_size = max(1, origin_batch_size) - infer_batch_size = triton.cdiv(infer_batch_size, self.decode_batch_alignment) * self.decode_batch_alignment + infer_batch_size = self._align_decode_batch_size(max(1, origin_batch_size)) # CUDA Graph 可能继续向上对齐 batch size,并因此加入 seq_len=2 的 # dummy request。先用最终可能出现的 KV 长度判断 graph,再统一 padding 一次。 @@ -825,8 +824,7 @@ def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1 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.decode_batch_alignment) * self.decode_batch_alignment + infer_batch_size = self._align_decode_batch_size(max(1, origin_batch_size0, origin_batch_size1)) 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) diff --git a/unit_tests/common/basemodel/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py index bf6e52db6e..4ce5329366 100644 --- a/unit_tests/common/basemodel/test_cuda_graph_layout.py +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -3,7 +3,6 @@ import pytest import lightllm.common.basemodel.cuda_graph as cuda_graph_module -import lightllm.common.basemodel.basemodel as basemodel_module import lightllm.common.basemodel.mtp_manager as mtp_manager_module from lightllm.common.basemodel.basemodel import TpPartBaseModel from lightllm.common.basemodel.cuda_graph import CudaGraph @@ -93,32 +92,20 @@ def test_mtp_tpsp_layout(monkeypatch, _graph_args, tp_size, mtp_step, dynamic, i args.mtp_mode = "eagle_with_att" args.mtp_step = mtp_step args.mtp_dynamic_verify = dynamic - args.mem_fraction = 0.8 - monkeypatch.setattr(basemodel_module, "get_env_start_args", lambda: args) monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) - monkeypatch.setattr(MtpManager, "_instance", None) - monkeypatch.setattr(basemodel_module, "get_dp_world_size", lambda: tp_size) - monkeypatch.setattr(basemodel_module, "get_llm_data_type", lambda: None) - monkeypatch.setattr(basemodel_module, "TorchMemorySaverWrapper", lambda _: None) - - class StopBeforeWeights(Exception): - pass - - def stop_init(self): - raise StopBeforeWeights - - monkeypatch.setattr(TpPartBaseModel, "_init_config", stop_init) model = TpPartBaseModel.__new__(TpPartBaseModel) + model.args = args + model.mtp_manager = MtpManager() + model.tp_world_size_ = tp_size + model.enable_tpsp_mix_mode = True model.is_mtp_draft_model = is_draft - with pytest.raises(StopBeforeWeights): - model.__init__(dict(run_mode="normal", weight_dir="unused", max_total_token_num=128, graph_max_batch_size=7)) + model._init_decode_batch_layout(max_requests=7) width = 1 if dynamic or is_draft else mtp_step + 1 assert model.decode_batch_alignment % width == 0 assert model.decode_batch_alignment % tp_size == 0 if width == 1: assert model.decode_batch_alignment == tp_size - assert model.max_req_num == 1000 # Physical padding does not change request capacity. assert model.graph_max_batch_size % model.decode_batch_alignment == 0 logical_max = 7 // 2 if overlap else 7 physical_max = logical_max * (1 if is_draft else mtp_step + 1) From a360f64f215be9ba35a860e0f6421d2fe9fe6fcb Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 21 Sep 2026 16:18:23 +0800 Subject: [PATCH 03/11] docs(mtp): clarify decode batch layout --- lightllm/common/basemodel/basemodel.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 829950d618..498023e1a2 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -267,14 +267,20 @@ def _init_att_backend1(self): return def _init_decode_batch_layout(self, max_requests: int): + # overlap decode 将请求拆成两个 microbatch,单个 CUDA Graph 只需覆盖其中一半。 if self.args.enable_decode_microbatch_overlap: max_requests //= 2 mtp = self.mtp_manager + # 固定 verify 的主模型每请求包含多个连续 MTP 行,动态 verify 会压缩为变长行数; + # grow step 分别保留这两种图捕获粒度。TP/SP 同时开启时还必须整除 TP, + # 因此用最小公倍数作为真实 batch 和 CUDA Graph 档位的统一对齐单位。 self.decode_batch_alignment = math.lcm( mtp.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model), self.tp_world_size_ if self.enable_tpsp_mix_mode else 1, ) + # graph 上限按物理 decode 行数计算:主模型 MTP verify 会扩展请求行, + # 各类 draft model 则由 MtpManager 给出自身的行数倍率。 max_rows = max_requests * mtp.get_decode_batch_multiplier(self.is_mtp_draft_model) self.graph_max_batch_size = self._align_decode_batch_size(max_rows) From 70c30d5694c3e8f663889392b06223ca25830780 Mon Sep 17 00:00:00 2001 From: sufubao Date: Mon, 21 Sep 2026 23:30:54 +0800 Subject: [PATCH 04/11] refactor(mtp): clarify decode batch size initialization --- lightllm/common/basemodel/basemodel.py | 8 ++++---- unit_tests/common/basemodel/test_cuda_graph_layout.py | 2 +- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 498023e1a2..5ddc3ca816 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -94,7 +94,7 @@ def __init__(self, kvargs): self.mem_fraction = kvargs.get("mem_fraction", 0.9) self.tp_world_size_ = get_dp_world_size() self.enable_tpsp_mix_mode = get_env_start_args().enable_tpsp_mix_mode - self._init_decode_batch_layout(kvargs.get("graph_max_batch_size", 16)) + self._init_decode_batch_sizes(kvargs.get("graph_max_batch_size", 16)) self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) self.prefill_graph: PrefillCudaGraph = None @@ -266,10 +266,10 @@ def _init_att_backend1(self): self.decode_att_backend1: BaseAttBackend = None return - def _init_decode_batch_layout(self, max_requests: int): + def _init_decode_batch_sizes(self, graph_max_requests: int): # overlap decode 将请求拆成两个 microbatch,单个 CUDA Graph 只需覆盖其中一半。 if self.args.enable_decode_microbatch_overlap: - max_requests //= 2 + graph_max_requests //= 2 mtp = self.mtp_manager # 固定 verify 的主模型每请求包含多个连续 MTP 行,动态 verify 会压缩为变长行数; @@ -281,7 +281,7 @@ def _init_decode_batch_layout(self, max_requests: int): ) # graph 上限按物理 decode 行数计算:主模型 MTP verify 会扩展请求行, # 各类 draft model 则由 MtpManager 给出自身的行数倍率。 - max_rows = max_requests * mtp.get_decode_batch_multiplier(self.is_mtp_draft_model) + max_rows = graph_max_requests * mtp.get_decode_batch_multiplier(self.is_mtp_draft_model) self.graph_max_batch_size = self._align_decode_batch_size(max_rows) def _align_decode_batch_size(self, batch_size: int) -> int: diff --git a/unit_tests/common/basemodel/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py index 4ce5329366..b1f2ecd4e7 100644 --- a/unit_tests/common/basemodel/test_cuda_graph_layout.py +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -99,7 +99,7 @@ def test_mtp_tpsp_layout(monkeypatch, _graph_args, tp_size, mtp_step, dynamic, i model.tp_world_size_ = tp_size model.enable_tpsp_mix_mode = True model.is_mtp_draft_model = is_draft - model._init_decode_batch_layout(max_requests=7) + model._init_decode_batch_sizes(graph_max_requests=7) width = 1 if dynamic or is_draft else mtp_step + 1 assert model.decode_batch_alignment % width == 0 From 7dfbc4261decaf8ad5b035caa7e26cfa7df3a275 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 22 Sep 2026 13:27:12 +0800 Subject: [PATCH 05/11] refactor(mtp): inline decode batch initialization --- lightllm/common/basemodel/basemodel.py | 30 ++++++------------- .../basemodel/test_cuda_graph_layout.py | 12 ++++---- .../common/basemodel/test_model_output.py | 6 ++-- 3 files changed, 18 insertions(+), 30 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 5ddc3ca816..b2baf2e2bb 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -94,7 +94,11 @@ def __init__(self, kvargs): self.mem_fraction = kvargs.get("mem_fraction", 0.9) self.tp_world_size_ = get_dp_world_size() self.enable_tpsp_mix_mode = get_env_start_args().enable_tpsp_mix_mode - self._init_decode_batch_sizes(kvargs.get("graph_max_batch_size", 16)) + self.graph_max_batch_size = kvargs.get("graph_max_batch_size", 16) + if self.args.enable_decode_microbatch_overlap: + self.graph_max_batch_size //= 2 + self.graph_max_batch_size *= self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) + self.graph_max_batch_size = self._align_decode_batch_size(self.graph_max_batch_size) self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) self.prefill_graph: PrefillCudaGraph = None @@ -266,27 +270,11 @@ def _init_att_backend1(self): self.decode_att_backend1: BaseAttBackend = None return - def _init_decode_batch_sizes(self, graph_max_requests: int): - # overlap decode 将请求拆成两个 microbatch,单个 CUDA Graph 只需覆盖其中一半。 - if self.args.enable_decode_microbatch_overlap: - graph_max_requests //= 2 - - mtp = self.mtp_manager - # 固定 verify 的主模型每请求包含多个连续 MTP 行,动态 verify 会压缩为变长行数; - # grow step 分别保留这两种图捕获粒度。TP/SP 同时开启时还必须整除 TP, - # 因此用最小公倍数作为真实 batch 和 CUDA Graph 档位的统一对齐单位。 - self.decode_batch_alignment = math.lcm( - mtp.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model), - self.tp_world_size_ if self.enable_tpsp_mix_mode else 1, - ) - # graph 上限按物理 decode 行数计算:主模型 MTP verify 会扩展请求行, - # 各类 draft model 则由 MtpManager 给出自身的行数倍率。 - max_rows = graph_max_requests * mtp.get_decode_batch_multiplier(self.is_mtp_draft_model) - self.graph_max_batch_size = self._align_decode_batch_size(max_rows) - def _align_decode_batch_size(self, batch_size: int) -> int: - # Preserve complete MTP verify groups and TP/SP shards, including dummy rows. - return triton.cdiv(batch_size, self.decode_batch_alignment) * self.decode_batch_alignment + alignment = self.mtp_manager.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model) + if self.args.enable_tpsp_mix_mode: + alignment = math.lcm(alignment, self.tp_world_size_) + return triton.cdiv(batch_size, alignment) * alignment def _init_cudagraph(self): # When graph covers the configured request length, it must also cover MTP's internal token margin. diff --git a/unit_tests/common/basemodel/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py index b1f2ecd4e7..ef1ee0bb4b 100644 --- a/unit_tests/common/basemodel/test_cuda_graph_layout.py +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -1,3 +1,4 @@ +import math from types import SimpleNamespace import pytest @@ -99,17 +100,14 @@ def test_mtp_tpsp_layout(monkeypatch, _graph_args, tp_size, mtp_step, dynamic, i model.tp_world_size_ = tp_size model.enable_tpsp_mix_mode = True model.is_mtp_draft_model = is_draft - model._init_decode_batch_sizes(graph_max_requests=7) width = 1 if dynamic or is_draft else mtp_step + 1 - assert model.decode_batch_alignment % width == 0 - assert model.decode_batch_alignment % tp_size == 0 - if width == 1: - assert model.decode_batch_alignment == tp_size - assert model.graph_max_batch_size % model.decode_batch_alignment == 0 logical_max = 7 // 2 if overlap else 7 physical_max = logical_max * (1 if is_draft else mtp_step + 1) - assert physical_max <= model.graph_max_batch_size < physical_max + model.decode_batch_alignment + model.graph_max_batch_size = model._align_decode_batch_size(physical_max) + alignment = math.lcm(width, tp_size) + assert model.graph_max_batch_size % alignment == 0 + assert physical_max <= model.graph_max_batch_size < physical_max + alignment sizes = CudaGraph.gen_cuda_graph_batch_sizes( batch_step_size_before_split=width, diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index b78e61812f..e1abec7898 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -116,7 +116,8 @@ def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): model = TpPartBaseModel.__new__(TpPartBaseModel) model.args = SimpleNamespace(enable_tpsp_mix_mode=enable_tpsp_mix_mode, page_size=1) model.tp_world_size_ = tp_world_size - model.decode_batch_alignment = tp_world_size if enable_tpsp_mix_mode else 1 + model.is_mtp_draft_model = False + model.mtp_manager = SimpleNamespace(get_decode_cuda_graph_grow_step_size=lambda _: 1) model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEXES=(99,), page_size=1) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) @@ -170,7 +171,8 @@ def test_fixed_mtp_decode_preserves_groups_and_real_rows(overlap, graph_mode, ro model = TpPartBaseModel.__new__(TpPartBaseModel) model.args = SimpleNamespace(enable_tpsp_mix_mode=True) model.tp_world_size_ = 8 - model.decode_batch_alignment = 24 # TP=8, mtp_step=2. + model.is_mtp_draft_model = False + model.mtp_manager = SimpleNamespace(get_decode_cuda_graph_grow_step_size=lambda _: 3) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88) seen_sizes = [] From fc6e4362356600d608b8912afb512bf052bfa442 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 22 Sep 2026 13:42:12 +0800 Subject: [PATCH 06/11] refactor(mtp): simplify decode layout policy --- lightllm/common/basemodel/mtp_manager.py | 32 ++++--------------- .../common/basemodel/test_mtp_manager.py | 16 ++++++---- 2 files changed, 15 insertions(+), 33 deletions(-) diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index be6c477b99..1acc5ef465 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -14,8 +14,6 @@ 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 @@ -28,26 +26,14 @@ 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.""" + """返回每请求的 decode 容量倍率;动态 verify 按未压缩的最大行数预留。""" 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 + return self.args.mtp_step + 1 # Block draft models decode mtp_step rows per logical request. if spec_mode in self._BLOCK_DRAFT_MODES: @@ -56,16 +42,10 @@ def get_decode_batch_multiplier(self, is_draft_model: bool) -> int: 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) + """返回 decode/graph 的基础对齐粒度;动态主模型压缩后允许任意行数。""" + if not is_draft_model and self.args.mtp_dynamic_verify: + return 1 + return self.get_decode_batch_multiplier(is_draft_model) def get_decode_draft_step(self, is_draft_model: bool) -> int: """Return the number of extra decode rows processed per request.""" diff --git a/unit_tests/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py index c7ad04f56b..dc833373eb 100644 --- a/unit_tests/common/basemodel/test_mtp_manager.py +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -49,17 +49,19 @@ def test_decode_batch_multiplier(monkeypatch, spec_mode, is_draft_model, expecte @pytest.mark.parametrize( - "dynamic_verify,is_draft_model,expected", + "spec_mode,dynamic_verify,is_draft_model,expected", [ - (False, False, 8), - (True, False, 1), - (False, True, 1), - (True, True, 1), + ("vanilla_with_att", False, False, 8), + ("vanilla_with_att", True, False, 1), + ("vanilla_with_att", False, True, 1), + ("vanilla_with_att", True, True, 1), + ("dspark", True, True, 7), + ("dflash", True, True, 7), ], ) -def test_decode_cuda_graph_grow_step_size(monkeypatch, dynamic_verify, is_draft_model, expected): +def test_decode_cuda_graph_grow_step_size(monkeypatch, spec_mode, dynamic_verify, is_draft_model, expected): args = SimpleNamespace( - mtp_mode="vanilla_with_att", + mtp_mode=spec_mode, mtp_step=7, mtp_dynamic_verify=dynamic_verify, ) From 9a048896acc5eccaa7c3fadcf1a783d16fbd3c7f Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 22 Sep 2026 13:43:27 +0800 Subject: [PATCH 07/11] refactor(mtp): rename decode batch alignment accessor --- lightllm/common/basemodel/basemodel.py | 4 ++-- lightllm/common/basemodel/mtp_manager.py | 2 +- unit_tests/common/basemodel/test_model_output.py | 4 ++-- unit_tests/common/basemodel/test_mtp_manager.py | 4 ++-- 4 files changed, 7 insertions(+), 7 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index b2baf2e2bb..400841d0db 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -271,7 +271,7 @@ def _init_att_backend1(self): return def _align_decode_batch_size(self, batch_size: int) -> int: - alignment = self.mtp_manager.get_decode_cuda_graph_grow_step_size(self.is_mtp_draft_model) + alignment = self.mtp_manager.get_decode_batch_alignment(self.is_mtp_draft_model) if self.args.enable_tpsp_mix_mode: alignment = math.lcm(alignment, self.tp_world_size_) return triton.cdiv(batch_size, alignment) * alignment @@ -282,7 +282,7 @@ def _init_cudagraph(self): self.graph_max_len_in_batch = max(self.graph_max_len_in_batch, self.max_seq_length) 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_grow_step_size = self.mtp_manager.get_decode_batch_alignment(self.is_mtp_draft_model) self.graph = ( None if self.disable_cudagraph diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index 1acc5ef465..6f1256a62e 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -41,7 +41,7 @@ def get_decode_batch_multiplier(self, is_draft_model: bool) -> int: return 1 - def get_decode_cuda_graph_grow_step_size(self, is_draft_model: bool) -> int: + def get_decode_batch_alignment(self, is_draft_model: bool) -> int: """返回 decode/graph 的基础对齐粒度;动态主模型压缩后允许任意行数。""" if not is_draft_model and self.args.mtp_dynamic_verify: return 1 diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index e1abec7898..a8a9da92db 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -117,7 +117,7 @@ def test_decode_pads_only_once_after_selecting_execution_path(monkeypatch): model.args = SimpleNamespace(enable_tpsp_mix_mode=enable_tpsp_mix_mode, page_size=1) model.tp_world_size_ = tp_world_size model.is_mtp_draft_model = False - model.mtp_manager = SimpleNamespace(get_decode_cuda_graph_grow_step_size=lambda _: 1) + model.mtp_manager = SimpleNamespace(get_decode_batch_alignment=lambda _: 1) model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEXES=(99,), page_size=1) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) @@ -172,7 +172,7 @@ def test_fixed_mtp_decode_preserves_groups_and_real_rows(overlap, graph_mode, ro model.args = SimpleNamespace(enable_tpsp_mix_mode=True) model.tp_world_size_ = 8 model.is_mtp_draft_model = False - model.mtp_manager = SimpleNamespace(get_decode_cuda_graph_grow_step_size=lambda _: 3) + model.mtp_manager = SimpleNamespace(get_decode_batch_alignment=lambda _: 3) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88) seen_sizes = [] diff --git a/unit_tests/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py index dc833373eb..0438089b08 100644 --- a/unit_tests/common/basemodel/test_mtp_manager.py +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -59,7 +59,7 @@ def test_decode_batch_multiplier(monkeypatch, spec_mode, is_draft_model, expecte ("dflash", True, True, 7), ], ) -def test_decode_cuda_graph_grow_step_size(monkeypatch, spec_mode, dynamic_verify, is_draft_model, expected): +def test_decode_batch_alignment(monkeypatch, spec_mode, dynamic_verify, is_draft_model, expected): args = SimpleNamespace( mtp_mode=spec_mode, mtp_step=7, @@ -67,7 +67,7 @@ def test_decode_cuda_graph_grow_step_size(monkeypatch, spec_mode, 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 + assert MtpManager.get_instance().get_decode_batch_alignment(is_draft_model) == expected @pytest.mark.parametrize( From a988d26cdff9ff07518d9ec5e4a4af3c96e7ea87 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 22 Sep 2026 13:47:11 +0800 Subject: [PATCH 08/11] refactor(mtp): name decode tokens per request explicitly --- lightllm/common/basemodel/basemodel.py | 6 +++--- lightllm/common/basemodel/mtp_manager.py | 8 ++++---- unit_tests/common/basemodel/test_mtp_manager.py | 8 ++++---- 3 files changed, 11 insertions(+), 11 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 400841d0db..60652a1fca 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -97,7 +97,7 @@ def __init__(self, kvargs): self.graph_max_batch_size = kvargs.get("graph_max_batch_size", 16) if self.args.enable_decode_microbatch_overlap: self.graph_max_batch_size //= 2 - self.graph_max_batch_size *= self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) + self.graph_max_batch_size *= self.mtp_manager.get_decode_tokens_per_request(self.is_mtp_draft_model) self.graph_max_batch_size = self._align_decode_batch_size(self.graph_max_batch_size) self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) @@ -281,14 +281,14 @@ def _init_cudagraph(self): if self.args.mtp_mode is not None and self.graph_max_len_in_batch >= self.args.max_req_total_len: self.graph_max_len_in_batch = max(self.graph_max_len_in_batch, self.max_seq_length) - decode_batch_multiplier = self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) + decode_tokens_per_request = self.mtp_manager.get_decode_tokens_per_request(self.is_mtp_draft_model) cuda_graph_grow_step_size = self.mtp_manager.get_decode_batch_alignment(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, + split_batch_size=self.args.graph_split_batch_size * decode_tokens_per_request, 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, diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index 6f1256a62e..34008f7ac9 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -25,8 +25,8 @@ def get_instance(cls) -> "MtpManager": def __init__(self): self.args = get_env_start_args() - def get_decode_batch_multiplier(self, is_draft_model: bool) -> int: - """返回每请求的 decode 容量倍率;动态 verify 按未压缩的最大行数预留。""" + def get_decode_tokens_per_request(self, is_draft_model: bool) -> int: + """返回每请求的 decode token 数;动态 verify 返回压缩前的最大 token 数,用于容量规划。""" spec_mode = self.args.mtp_mode if spec_mode is None: @@ -45,12 +45,12 @@ def get_decode_batch_alignment(self, is_draft_model: bool) -> int: """返回 decode/graph 的基础对齐粒度;动态主模型压缩后允许任意行数。""" if not is_draft_model and self.args.mtp_dynamic_verify: return 1 - return self.get_decode_batch_multiplier(is_draft_model) + return self.get_decode_tokens_per_request(is_draft_model) 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 + return self.get_decode_tokens_per_request(is_draft_model) - 1 def create_hidden_collector( self, diff --git a/unit_tests/common/basemodel/test_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py index 0438089b08..3ab2c926db 100644 --- a/unit_tests/common/basemodel/test_mtp_manager.py +++ b/unit_tests/common/basemodel/test_mtp_manager.py @@ -20,14 +20,14 @@ def _reset_mtp_manager(): MtpManager._instance = None -def _decode_batch_multiplier(monkeypatch, spec_mode, *, is_draft_model, mtp_step=7): +def _decode_tokens_per_request(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) + return MtpManager.get_instance().get_decode_tokens_per_request(is_draft_model) @pytest.mark.parametrize( @@ -44,8 +44,8 @@ def _decode_batch_multiplier(monkeypatch, spec_mode, *, is_draft_model, mtp_step ("dflash", True, 7), ], ) -def test_decode_batch_multiplier(monkeypatch, spec_mode, is_draft_model, expected): - assert _decode_batch_multiplier(monkeypatch, spec_mode, is_draft_model=is_draft_model) == expected +def test_decode_tokens_per_request(monkeypatch, spec_mode, is_draft_model, expected): + assert _decode_tokens_per_request(monkeypatch, spec_mode, is_draft_model=is_draft_model) == expected @pytest.mark.parametrize( From a3f6002eb68a8b205a2c23244ed218ca0d94409c Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 22 Sep 2026 13:59:31 +0800 Subject: [PATCH 09/11] test(mtp): cover constructor capacity and dynamic decode padding --- .../basemodel/test_cuda_graph_layout.py | 21 ++++++++++++++----- .../common/basemodel/test_model_output.py | 19 ++++++++++++----- .../common/basemodel/test_overlap_utils.py | 1 + 3 files changed, 31 insertions(+), 10 deletions(-) diff --git a/unit_tests/common/basemodel/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py index ef1ee0bb4b..30d4c54368 100644 --- a/unit_tests/common/basemodel/test_cuda_graph_layout.py +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -3,6 +3,7 @@ import pytest +import lightllm.common.basemodel.basemodel as basemodel_module import lightllm.common.basemodel.cuda_graph as cuda_graph_module import lightllm.common.basemodel.mtp_manager as mtp_manager_module from lightllm.common.basemodel.basemodel import TpPartBaseModel @@ -95,16 +96,26 @@ def test_mtp_tpsp_layout(monkeypatch, _graph_args, tp_size, mtp_step, dynamic, i args.mtp_dynamic_verify = dynamic monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) model = TpPartBaseModel.__new__(TpPartBaseModel) - model.args = args - model.mtp_manager = MtpManager() - model.tp_world_size_ = tp_size - model.enable_tpsp_mix_mode = True model.is_mtp_draft_model = is_draft + monkeypatch.setattr(basemodel_module, "get_env_start_args", lambda: args) + monkeypatch.setattr(basemodel_module, "get_llm_data_type", lambda: None) + monkeypatch.setattr(basemodel_module, "get_dp_world_size", lambda: tp_size) + monkeypatch.setattr(MtpManager, "_instance", MtpManager()) + + class StopBeforeWeights(Exception): + pass + + def stop_before_weights(): + raise StopBeforeWeights + + # 执行真实构造函数的容量计算,在读取模型配置、分配 GPU 权重前停止。 + monkeypatch.setattr(model, "_init_config", stop_before_weights) + with pytest.raises(StopBeforeWeights): + model.__init__(dict(run_mode="normal", weight_dir="", max_total_token_num=1024, graph_max_batch_size=7)) width = 1 if dynamic or is_draft else mtp_step + 1 logical_max = 7 // 2 if overlap else 7 physical_max = logical_max * (1 if is_draft else mtp_step + 1) - model.graph_max_batch_size = model._align_decode_batch_size(physical_max) alignment = math.lcm(width, tp_size) assert model.graph_max_batch_size % alignment == 0 assert physical_max <= model.graph_max_batch_size < physical_max + alignment diff --git a/unit_tests/common/basemodel/test_model_output.py b/unit_tests/common/basemodel/test_model_output.py index a8a9da92db..3f6664d031 100644 --- a/unit_tests/common/basemodel/test_model_output.py +++ b/unit_tests/common/basemodel/test_model_output.py @@ -165,20 +165,22 @@ def create_infer_state(model_input): @pytest.mark.parametrize("overlap", [False, True]) +@pytest.mark.parametrize("dynamic", [False, True]) @pytest.mark.parametrize("graph_mode", ["eager", "capture", "replay"]) @pytest.mark.parametrize("rows", [0, 3, 9, 24, 27]) -def test_fixed_mtp_decode_preserves_groups_and_real_rows(overlap, graph_mode, rows): +def test_mtp_decode_preserves_groups_and_real_rows(overlap, dynamic, graph_mode, rows): model = TpPartBaseModel.__new__(TpPartBaseModel) model.args = SimpleNamespace(enable_tpsp_mix_mode=True) model.tp_world_size_ = 8 model.is_mtp_draft_model = False - model.mtp_manager = SimpleNamespace(get_decode_batch_alignment=lambda _: 3) + model.mtp_manager = SimpleNamespace(get_decode_batch_alignment=lambda _: 1 if dynamic else 3) model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88) seen_sizes = [] def create_state(model_input, microbatch_index=0): def check_layout(): - assert model_input.batch_size % 3 == 0 + if not dynamic: + assert model_input.batch_size % 3 == 0 assert model_input.batch_size % 8 == 0 seen_sizes.append(model_input.batch_size) @@ -199,7 +201,7 @@ def forward(state): if graph_mode != "eager": model.graph = SimpleNamespace( can_run=lambda batch_size, max_len_in_batch: batch_size <= 24, - find_closest_graph_batch_size=lambda batch_size: 24, + find_closest_graph_batch_size=lambda batch_size: ((batch_size + 7) // 8 * 8 if dynamic else 24), need_capture=lambda batch_size: graph_mode == "capture", capture_decode=lambda func, state, **kwargs: func(state, **kwargs), replay=lambda state, infer_state1=None: ( @@ -210,6 +212,12 @@ def forward(state): model_input = model._create_padded_decode_model_input(_create_empty_decode_input(), rows) model_input.b_req_idx = torch.arange(rows, dtype=torch.int32) // 3 model_input.b_seq_len = 10 + torch.arange(rows, dtype=torch.int32) % 3 + if dynamic: + # 压缩后每请求分别保留 1、2、3 个 token,保持各请求内部的 token 顺序。 + request_rows = [(req, step) for req in range(rows) for step in range(req % 3 + 1)][:rows] + model_input.b_req_idx = torch.tensor([req for req, _ in request_rows], dtype=torch.int32) + model_input.b_mtp_index = torch.tensor([step for _, step in request_rows], dtype=torch.int32) + model_input.b_seq_len = 10 + model_input.b_mtp_index expected = torch.stack((model_input.b_req_idx, model_input.b_seq_len), dim=1) if overlap: output, empty_output = model._microbatch_overlap_decode_cuda(model_input, _create_empty_decode_input()) @@ -217,6 +225,7 @@ def forward(state): else: output = model._decode(model_input) - assert seen_sizes == ([48 if rows > 24 else 24] * (2 if overlap else 1)) + expected_size = (max(1, rows) + 7) // 8 * 8 if dynamic else (48 if rows > 24 else 24) + assert seen_sizes == [expected_size] * (2 if overlap else 1) torch.testing.assert_close(output.logits, expected) torch.testing.assert_close(torch.stack((model_input.b_req_idx, model_input.b_seq_len), dim=1), expected) diff --git a/unit_tests/common/basemodel/test_overlap_utils.py b/unit_tests/common/basemodel/test_overlap_utils.py index 6e71bfd338..54f2ccacf1 100644 --- a/unit_tests/common/basemodel/test_overlap_utils.py +++ b/unit_tests/common/basemodel/test_overlap_utils.py @@ -181,6 +181,7 @@ def test_overlap_decode_cuda_pads_empty_side_and_unpads_outputs(monkeypatch): model = TpPartBaseModel.__new__(TpPartBaseModel) model.args = SimpleNamespace(enable_tpsp_mix_mode=True, page_size=1) model.tp_world_size_ = 2 + model.mtp_manager = SimpleNamespace(get_decode_batch_alignment=lambda _: 1) model.graph = None model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=88, req_to_token_indexs=object()) model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEXES=(77,), page_size=1) From 153b1b482605d79084fd402d3b2d5f7180b5ba57 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 22 Sep 2026 14:03:11 +0800 Subject: [PATCH 10/11] refactor(mtp): preserve graph capacity initialization order --- lightllm/common/basemodel/basemodel.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 60652a1fca..5a2d5b77d8 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -84,7 +84,11 @@ def __init__(self, kvargs): self.max_seq_length = kvargs.get("max_seq_length", 1024 * 5) self.return_all_prompt_logics = kvargs.get("return_all_prompt_logics", False) self.data_type = get_llm_data_type() + self.graph_max_batch_size = kvargs.get("graph_max_batch_size", 16) + if self.args.enable_decode_microbatch_overlap: + self.graph_max_batch_size //= 2 self.mtp_manager = MtpManager.get_instance() + self.graph_max_batch_size *= self.mtp_manager.get_decode_tokens_per_request(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) @@ -94,10 +98,6 @@ def __init__(self, kvargs): self.mem_fraction = kvargs.get("mem_fraction", 0.9) self.tp_world_size_ = get_dp_world_size() self.enable_tpsp_mix_mode = get_env_start_args().enable_tpsp_mix_mode - self.graph_max_batch_size = kvargs.get("graph_max_batch_size", 16) - if self.args.enable_decode_microbatch_overlap: - self.graph_max_batch_size //= 2 - self.graph_max_batch_size *= self.mtp_manager.get_decode_tokens_per_request(self.is_mtp_draft_model) self.graph_max_batch_size = self._align_decode_batch_size(self.graph_max_batch_size) self.torch_memory_saver = TorchMemorySaverWrapper(self.args.enable_torch_memory_saver) From cd7efee833a532bcdc212c5f4c961fa9bdd3a621 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 22 Sep 2026 14:52:05 +0800 Subject: [PATCH 11/11] refactor(mtp): restore explicit draft mode branches --- lightllm/common/basemodel/mtp_manager.py | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index 34008f7ac9..d31b33b3d8 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -14,6 +14,8 @@ 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 @@ -32,8 +34,20 @@ def get_decode_tokens_per_request(self, is_draft_model: bool) -> int: 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 self.args.mtp_step + 1 + 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: