diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 7d1551e05..5a2d5b77d 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 @@ -84,15 +85,10 @@ def __init__(self, kvargs): 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 - ) + 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.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_len_in_batch = kvargs.get("graph_max_len_in_batch", 8192) self.disable_cudagraph = kvargs.get("disable_cudagraph", False) @@ -102,6 +98,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.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 @@ -273,19 +270,25 @@ def _init_att_backend1(self): self.decode_att_backend1: BaseAttBackend = None return + def _align_decode_batch_size(self, batch_size: int) -> int: + 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 + 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: 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) + 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, @@ -592,11 +595,8 @@ def _decode( ) 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(infer_batch_size, self.tp_world_size_) * self.tp_world_size_ + # 空 DP rank 也需要完整的 dummy verify group;TP/SP 切分不能破坏 MTP 分组。 + 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 一次。 @@ -818,8 +818,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.tp_world_size_) * self.tp_world_size_ + 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/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index dd15af99f..ee888c1a0 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/lightllm/common/basemodel/mtp_manager.py b/lightllm/common/basemodel/mtp_manager.py index be6c477b9..d31b33b3d 100644 --- a/lightllm/common/basemodel/mtp_manager.py +++ b/lightllm/common/basemodel/mtp_manager.py @@ -27,8 +27,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: - """Return the physical decode rows used by one logical request.""" + 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: @@ -55,22 +55,16 @@ 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) + 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_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_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py index 17ddcd98d..30d4c5436 100644 --- a/unit_tests/common/basemodel/test_cuda_graph_layout.py +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -1,9 +1,14 @@ +import math from types import SimpleNamespace 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 from lightllm.common.basemodel.cuda_graph import CudaGraph +from lightllm.common.basemodel.mtp_manager import MtpManager @pytest.fixture(autouse=True) @@ -76,3 +81,53 @@ 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 + monkeypatch.setattr(mtp_manager_module, "get_env_start_args", lambda: args) + model = TpPartBaseModel.__new__(TpPartBaseModel) + 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) + 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, + 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 4b723e6a1..3f6664d03 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,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.is_mtp_draft_model = False + 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()) @@ -159,3 +162,70 @@ 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("dynamic", [False, True]) +@pytest.mark.parametrize("graph_mode", ["eager", "capture", "replay"]) +@pytest.mark.parametrize("rows", [0, 3, 9, 24, 27]) +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 _: 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(): + if not dynamic: + 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: ((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: ( + 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 + 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()) + assert empty_output.logits.shape == (0, 2) + else: + output = model._decode(model_input) + + 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_mtp_manager.py b/unit_tests/common/basemodel/test_mtp_manager.py index c7ad04f56..3ab2c926d 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,28 +44,30 @@ 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( - "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_batch_alignment(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, ) 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( diff --git a/unit_tests/common/basemodel/test_overlap_utils.py b/unit_tests/common/basemodel/test_overlap_utils.py index 6e71bfd33..54f2ccacf 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)