Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 17 additions & 18 deletions lightllm/common/basemodel/basemodel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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 一次。
Expand Down Expand Up @@ -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)
Expand Down
5 changes: 4 additions & 1 deletion lightllm/common/basemodel/cuda_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down
22 changes: 8 additions & 14 deletions lightllm/common/basemodel/mtp_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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,
Expand Down
55 changes: 55 additions & 0 deletions unit_tests/common/basemodel/test_cuda_graph_layout.py
Original file line number Diff line number Diff line change
@@ -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)
Expand Down Expand Up @@ -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]
70 changes: 70 additions & 0 deletions unit_tests/common/basemodel/test_model_output.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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())

Expand Down Expand Up @@ -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)
26 changes: 14 additions & 12 deletions unit_tests/common/basemodel/test_mtp_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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(
Expand Down
1 change: 1 addition & 0 deletions unit_tests/common/basemodel/test_overlap_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading